diff --git a/src/cdtools/reconstructors/base.py b/src/cdtools/reconstructors/base.py index 22067ea..45c806c 100644 --- a/src/cdtools/reconstructors/base.py +++ b/src/cdtools/reconstructors/base.py @@ -218,12 +218,15 @@ class Reconstructor: return total_loss # This takes the step for this minibatch - loss += self.optimizer.step(closure).detach().cpu().numpy() + loss += self.optimizer.step(closure).detach() if hasattr(self.model, 'loss_normalizer') and \ self.model.loss_normalizer is not None: loss = self.model.loss_normalizer.normalize_loss(loss) + # Make sure to return a scalar value which is fully numpy + loss = loss.cpu().numpy()[()] + # We step the scheduler after the full epoch if self.scheduler is not None: self.scheduler.step(loss) diff --git a/tests/test_reconstructors.py b/tests/test_reconstructors.py index c9cade1..378330e 100644 --- a/tests/test_reconstructors.py +++ b/tests/test_reconstructors.py @@ -1,3 +1,4 @@ +from numbers import Number import pytest import time import cdtools @@ -114,6 +115,11 @@ def test_Adam_gold_balls(gold_ball_cxi, reconstruction_device, show_plot): time.sleep(3) plt.close('all') + + # Check that the losses returned in loss_history are not torch tensors + assert isinstance(model.loss_history[-1], Number) and \ + not isinstance(model.loss_history[-1], t.Tensor) + # Ensure equivalency between the model reconstructions during the first # pass, where they should be identical assert np.allclose(model_recon.loss_history[:epoch_tup[0]], model.loss_history[:epoch_tup[0]])