linting test_simple_ptycho.py

This commit is contained in:
gnzng
2025-07-07 13:59:35 -07:00
parent 745f7b73e6
commit ecb8ee4864
+6 -5
View File
@@ -1,17 +1,18 @@
import pytest
import cdtools
from matplotlib import pyplot as plt
import torch as t
import cdtools
# Force all reconstructions to use the same RNG seed
t.manual_seed(0)
@pytest.mark.slow
def test_simple_ptycho(lab_ptycho_cxi, reconstruction_device, show_plot):
dataset = cdtools.datasets.Ptycho2DDataset.from_cxi(lab_ptycho_cxi)
model = cdtools.models.SimplePtycho.from_dataset(dataset)
model.to(device=reconstruction_device)
dataset.get_as(device=reconstruction_device)
@@ -19,10 +20,10 @@ def test_simple_ptycho(lab_ptycho_cxi, reconstruction_device, show_plot):
print(model.report())
if show_plot and model.epoch % 10 == 0:
model.inspect(dataset)
if show_plot:
model.inspect(dataset)
model.compare(dataset)
# If this fails, the reconstruction got worse
assert model.loss_history[-1] < 0.013