Skip to content

Commit be8f59a

Browse files
committed
black
1 parent 7b18031 commit be8f59a

15 files changed

Lines changed: 87 additions & 116 deletions

File tree

deepinv_bench/__init__.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1 @@
1-
from .run import run_benchmark
1+
from .run import run_benchmark

deepinv_bench/base/runner_solver.py

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -13,14 +13,13 @@ class Solver(BaseSolver):
1313
To use this class, set the `model` and `name` attributes to run
1414
it in the model.
1515
"""
16+
1617
name = "runner-solver"
1718
model = None
1819
parameters = {}
1920

2021
def set_objective(self, train_dataset=None, physics=None):
21-
device = (
22-
dinv.utils.get_freer_gpu() if torch.cuda.is_available() else "cpu"
23-
)
22+
device = dinv.utils.get_freer_gpu() if torch.cuda.is_available() else "cpu"
2423
# Make sure the model is on the correct device and has a device
2524
# attribute to we can properly move the test data later.
2625
self.model.to(device)

deepinv_bench/benchmark_template/datasets/dataset.py

Lines changed: 8 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -14,29 +14,24 @@ class Dataset(BaseDataset):
1414
# IMPORTANT: the names of physics and noise should match
1515
# the ones defined in deepinv.physics with exact same spelling
1616
parameters = {
17-
'physics': ['physics_name'],
18-
'noise': ['noise_name'],
19-
'img_size': [256],
17+
"physics": ["physics_name"],
18+
"noise": ["noise_name"],
19+
"img_size": [256],
2020
# add any other parameter you might need
2121
}
2222

2323
def get_data(self):
2424
root = get_data_path("Set14_HR")
2525

26-
transform = transforms.Compose([
27-
transforms.Resize((self.img_size, self.img_size)),
28-
transforms.ToTensor()
29-
])
26+
transform = transforms.Compose(
27+
[transforms.Resize((self.img_size, self.img_size)), transforms.ToTensor()]
28+
)
3029

3130
# load the dataset
32-
dataset = dinv.datasets.Set14HR(
33-
root, download=True, transform=transform
34-
)
31+
dataset = dinv.datasets.Set14HR(root, download=True, transform=transform)
3532

3633
# define the physics according to the parameters
37-
physics = dinv.physics.Denoising(
38-
noise_model=dinv.physics.GaussianNoise()
39-
)
34+
physics = dinv.physics.Denoising(noise_model=dinv.physics.GaussianNoise())
4035

4136
return dict(
4237
dataset=dataset,

deepinv_bench/benchmark_template/objective.py

Lines changed: 6 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -17,21 +17,18 @@ class Objective(BaseObjective):
1717
# Bump it up if the benchmark depends on a new feature of benchopt.
1818
min_benchopt_version = "1.8"
1919

20-
sampling_strategy = 'run_once'
20+
sampling_strategy = "run_once"
2121

2222
def set_data(self, dataset, physics):
2323
self.dataset = dataset
2424
self.physics = physics
2525

2626
def evaluate_result(self, model):
27-
device = getattr(model, 'device', None)
27+
device = getattr(model, "device", None)
2828
self.physics = self.physics.to(device)
2929

3030
# change metrics if needed
31-
metrics = [
32-
dinv.loss.PSNR(),
33-
dinv.loss.NIQE(device=device)
34-
]
31+
metrics = [dinv.loss.PSNR(), dinv.loss.NIQE(device=device)]
3532

3633
results = dinv.test(
3734
model,
@@ -40,14 +37,15 @@ def evaluate_result(self, model):
4037
online_measurements=True,
4138
device=device,
4239
metrics=metrics,
43-
compare_no_learning=False
40+
compare_no_learning=False,
4441
)
4542

4643
return results
4744

4845
def get_one_result(self):
4946
class DummyModel:
50-
def eval(self): pass
47+
def eval(self):
48+
pass
5149

5250
def __call__(self, x, physics=None):
5351
return physics.A_adjoint(x)

deepinv_bench/benchmark_template/solvers/solver1.py

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -5,15 +5,13 @@
55

66

77
class Solver(BaseSolver):
8-
name = 'solver_name'
8+
name = "solver_name"
99

1010
# add any hyper-parameters here
1111
parameters = {}
1212

1313
def set_objective(self, train_dataset=None, physics=None):
14-
device = (
15-
dinv.utils.get_freer_gpu() if torch.cuda.is_available() else "cpu"
16-
)
14+
device = dinv.utils.get_freer_gpu() if torch.cuda.is_available() else "cpu"
1715

1816
# replace by your model, should take (y, physics) as input
1917
self.model = dinv.models.DnCNN(device=device)

deepinv_bench/cbsd500_gaussian_denoising/datasets/cbsd68.py

Lines changed: 9 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -14,27 +14,22 @@ class Dataset(BaseDataset):
1414

1515
name = "CBSD68"
1616
parameters = {
17-
'physics': ['Denoising'],
18-
'noise': ['GaussianNoise'],
19-
'sigma': [0.1],
20-
'img_size': [256],
21-
'debug': [False],
17+
"physics": ["Denoising"],
18+
"noise": ["GaussianNoise"],
19+
"sigma": [0.1],
20+
"img_size": [256],
21+
"debug": [False],
2222
}
2323

24-
test_parameters = {
25-
"debug": [True]
26-
}
24+
test_parameters = {"debug": [True]}
2725

2826
def get_data(self):
2927

3028
root = get_data_path("CBSD68")
31-
transform = transforms.Compose([
32-
transforms.Resize((self.img_size, self.img_size)),
33-
transforms.ToTensor()
34-
])
35-
dataset = CBSD68(
36-
root, download=True, transform=transform
29+
transform = transforms.Compose(
30+
[transforms.Resize((self.img_size, self.img_size)), transforms.ToTensor()]
3731
)
32+
dataset = CBSD68(root, download=True, transform=transform)
3833

3934
if self.debug:
4035
dataset = dinv.torch.utils.data.Subset(dataset, [0, 1, 2])

deepinv_bench/cbsd500_gaussian_denoising/objective.py

Lines changed: 5 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -25,13 +25,10 @@ def set_data(self, dataset, physics):
2525
self.physics = physics
2626

2727
def evaluate_result(self, model):
28-
device = getattr(model, 'device', None)
28+
device = getattr(model, "device", None)
2929
self.physics = self.physics.to(device)
3030

31-
metrics = [
32-
dinv.loss.PSNR(),
33-
dinv.loss.NIQE(device=device)
34-
]
31+
metrics = [dinv.loss.PSNR(), dinv.loss.NIQE(device=device)]
3532

3633
results = dinv.test(
3734
model,
@@ -40,15 +37,16 @@ def evaluate_result(self, model):
4037
online_measurements=True,
4138
device=device,
4239
metrics=metrics,
43-
compare_no_learning=False
40+
compare_no_learning=False,
4441
)
4542

4643
return results
4744

4845
def get_one_result(self):
4946

5047
class DummyModel:
51-
def eval(self): pass
48+
def eval(self):
49+
pass
5250

5351
def __call__(self, x, physics=None):
5452
return physics.A_adjoint(x)

deepinv_bench/cbsd500_gaussian_denoising/solvers/drunet.py

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -4,14 +4,12 @@
44

55

66
class Solver(BaseSolver):
7-
name = 'DRUNet'
7+
name = "DRUNet"
88

99
parameters = {}
1010

1111
def set_objective(self, train_dataset=None, physics=None):
12-
self.model = dinv.models.ArtifactRemoval(
13-
dinv.models.DRUNet()
14-
)
12+
self.model = dinv.models.ArtifactRemoval(dinv.models.DRUNet())
1513

1614
def run(self, _):
1715
pass

deepinv_bench/cbsd500_gaussian_denoising/solvers/restformer.py

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -4,14 +4,12 @@
44

55

66
class Solver(BaseSolver):
7-
name = 'Restormer'
7+
name = "Restormer"
88

99
parameters = {}
1010

1111
def set_objective(self, train_dataset=None, physics=None):
12-
self.model = dinv.models.ArtifactRemoval(
13-
dinv.models.Restormer()
14-
)
12+
self.model = dinv.models.ArtifactRemoval(dinv.models.Restormer())
1513

1614
def run(self, _):
1715
pass

deepinv_bench/div2k_gaussian_deblurring/datasets/div2k.py

Lines changed: 10 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -14,33 +14,28 @@ class Dataset(BaseDataset):
1414

1515
name = "DIV2K"
1616
parameters = {
17-
'physics': ['Blur'],
18-
'noise': ['GaussianNoise'],
19-
'sigma': [0.1],
20-
'img_size': [256],
21-
'debug': [False],
17+
"physics": ["Blur"],
18+
"noise": ["GaussianNoise"],
19+
"sigma": [0.1],
20+
"img_size": [256],
21+
"debug": [False],
2222
}
2323

24-
test_parameters = {
25-
"debug": [True]
26-
}
24+
test_parameters = {"debug": [True]}
2725

2826
def get_data(self):
2927
root = get_data_path("DIV2K")
30-
transform = transforms.Compose([
31-
transforms.Resize((self.img_size, self.img_size)),
32-
transforms.ToTensor()
33-
])
34-
dataset = DIV2K(
35-
root, mode="val", download=True, transform=transform
28+
transform = transforms.Compose(
29+
[transforms.Resize((self.img_size, self.img_size)), transforms.ToTensor()]
3630
)
31+
dataset = DIV2K(root, mode="val", download=True, transform=transform)
3732
if self.debug:
3833
dataset = dinv.torch.utils.data.Subset(dataset, [0, 1, 2])
3934

4035
return dict(
4136
dataset=dataset,
4237
physics=Blur(
4338
filter=dinv.physics.blur.gaussian_blur(2),
44-
noise_model=GaussianNoise(0.05)
39+
noise_model=GaussianNoise(0.05),
4540
),
4641
)

0 commit comments

Comments
 (0)