mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-10-11 02:30:25 +02:00
Fix a bug where the normalizers stopped the loss history from being saved as a numpy scalar, and added test coverage
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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]])
|
||||
|
||||
Reference in New Issue
Block a user