diff --git a/LION/optimizers/EquivariantSolver.py b/LION/optimizers/EquivariantSolver.py index 9b6234b3..98c70d5a 100644 --- a/LION/optimizers/EquivariantSolver.py +++ b/LION/optimizers/EquivariantSolver.py @@ -46,7 +46,9 @@ def rotation_group(cardinality: int): assert 360 % cardinality == 0 angle_increment = 360 / cardinality - return [lambda x: TF.rotate(x, i * angle_increment) for i in range(cardinality)] + return [ + lambda x, i=i: TF.rotate(x, i * angle_increment) for i in range(cardinality) + ] @staticmethod def default_parameters() -> LIONParameter: diff --git a/LION/optimizers/GaussianDenoiserSolver.py b/LION/optimizers/GaussianDenoiserSolver.py index 51eb53ab..fc5ebcf9 100644 --- a/LION/optimizers/GaussianDenoiserSolver.py +++ b/LION/optimizers/GaussianDenoiserSolver.py @@ -156,7 +156,7 @@ def validate(self): outputs = self.model(y) validation_loss = np.append( validation_loss, - self.validation_fn(y.to(self.device), outputs.to(self.device)) + self.validation_fn(outputs.to(self.device), data.to(self.device)) .cpu() .numpy(), ) diff --git a/LION/optimizers/LIONsolver.py b/LION/optimizers/LIONsolver.py index 6ff5f641..22fa8a54 100644 --- a/LION/optimizers/LIONsolver.py +++ b/LION/optimizers/LIONsolver.py @@ -491,7 +491,7 @@ def __check_attribute( warnings.warn(f"Attribute {attr} is not callable") return 2 # just standrad type chekcking, error or warn, depends of settings - elif isinstance(getattr(self, attr), type): + elif not isinstance(getattr(self, attr), expected_type): if error: raise ValueError( f"Attribute {attr} is not of type {expected_type}, its {type(getattr(self, attr))}" @@ -599,7 +599,8 @@ def test(self): data = fdk(data, self.op) output = self.model(data.to(self.device)) test_loss = np.append( - test_loss, self.testing_fn(output, target.to(self.device)) + test_loss, + self.testing_fn(output, target.to(self.device)).cpu().numpy(), ) if self.verbose: