Add first tests for models, and replace t.tensor with t.as_tensor in models

This commit is contained in:
2024-08-30 16:10:26 +02:00
parent 21b618d2e0
commit 30a989fe3e
10 changed files with 269 additions and 105 deletions
+3 -3
View File
@@ -117,18 +117,18 @@ class CDIModel(t.nn.Module):
The datatype to convert the values to before registering
"""
self.register_buffer('det_basis',
t.tensor(detector_geometry['basis'],
t.as_tensor(detector_geometry['basis'],
dtype=dtype))
if 'distance' in detector_geometry \
and detector_geometry['distance'] is not None:
self.register_buffer('det_distance',
t.tensor(detector_geometry['distance'],
t.as_tensor(detector_geometry['distance'],
dtype=dtype))
if 'corner' in detector_geometry \
and detector_geometry['corner'] is not None:
self.register_buffer('det_corner',
t.tensor(detector_geometry['corner'],
t.as_tensor(detector_geometry['corner'],
dtype=dtype))
def get_detector_geometry(self):
+21 -21
View File
@@ -92,23 +92,23 @@ class Bragg2DPtycho(CDIModel):
super(Bragg2DPtycho, self).__init__()
self.register_buffer('wavelength',
t.tensor(wavelength, dtype=dtype))
t.as_tensor(wavelength, dtype=dtype))
self.store_detector_geometry(detector_geometry,
dtype=dtype)
self.register_buffer('min_translation',
t.tensor(min_translation, dtype=dtype))
t.as_tensor(min_translation, dtype=dtype))
self.register_buffer('median_propagation',
t.tensor(median_propagation, dtype=dtype))
t.as_tensor(median_propagation, dtype=dtype))
self.register_buffer('obj_basis',
t.tensor(obj_basis, dtype=dtype))
t.as_tensor(obj_basis, dtype=dtype))
if probe_basis is None:
self.register_buffer('probe_basis',
t.tensor(obj_basis, dtype=dtype))
t.as_tensor(obj_basis, dtype=dtype))
else:
self.register_buffer('probe_basis',
t.tensor(probe_basis, dtype=dtype))
t.as_tensor(probe_basis, dtype=dtype))
self.units = units
@@ -119,22 +119,22 @@ class Bragg2DPtycho(CDIModel):
np.array(obj_basis)[:,0])
surface_normal /= np.linalg.norm(surface_normal)
self.register_buffer('surface_normal',
t.tensor(surface_normal, dtype=dtype))
t.as_tensor(surface_normal, dtype=dtype))
if saturation is None:
self.saturation = None
else:
self.register_buffer('saturation',
t.tensor(saturation, dtype=dtype))
t.as_tensor(saturation, dtype=dtype))
if mask is None:
self.mask = None
else:
self.register_buffer('mask',
t.tensor(mask, dtype=t.bool))
t.as_tensor(mask, dtype=t.bool))
probe_guess = t.tensor(probe_guess, dtype=t.complex64)
obj_guess = t.tensor(obj_guess, dtype=t.complex64)
probe_guess = t.as_tensor(probe_guess, dtype=t.complex64)
obj_guess = t.as_tensor(obj_guess, dtype=t.complex64)
# We rescale the probe here so it learns at the same rate as the
# object
@@ -150,7 +150,7 @@ class Bragg2DPtycho(CDIModel):
if probe_support is None:
probe_support = t.ones_like(self.probe[0], dtype=t.bool)
self.register_buffer('probe_support',
t.tensor(probe_support, dtype=t.bool))
t.as_tensor(probe_support, dtype=t.bool))
self.probe.data *= self.probe_support
if background is None:
@@ -164,26 +164,26 @@ class Bragg2DPtycho(CDIModel):
self.weights = None
else:
# No incoherent + unstable here yet
self.weights = t.nn.Parameter(t.tensor(weights,
self.weights = t.nn.Parameter(t.as_tensor(weights,
dtype=t.float32))
if translation_offsets is None:
self.translation_offsets = None
else:
t_o = t.tensor(translation_offsets, dtype=t.float32)
t_o = t.as_tensor(translation_offsets, dtype=t.float32)
t_o = t_o / translation_scale
self.translation_offsets = t.nn.Parameter(t_o)
self.register_buffer('translation_scale',
t.tensor(translation_scale, dtype=dtype))
t.as_tensor(translation_scale, dtype=dtype))
self.register_buffer('oversampling',
t.tensor(oversampling, dtype=int))
t.as_tensor(oversampling, dtype=int))
self.register_buffer('propagate_probe',
t.tensor(propagate_probe, dtype=bool))
t.as_tensor(propagate_probe, dtype=bool))
self.register_buffer('correct_tilt',
t.tensor(correct_tilt, dtype=bool))
t.as_tensor(correct_tilt, dtype=bool))
if correct_tilt:
k_map, intensity_map = \
@@ -195,9 +195,9 @@ class Bragg2DPtycho(CDIModel):
self.wavelength,dtype=t.float32,
lens=lens)
self.register_buffer('k_map',
t.tensor(k_map, dtype=dtype))
t.as_tensor(k_map, dtype=dtype))
self.register_buffer('intensity_map',
t.tensor(intensity_map, dtype=dtype))
t.as_tensor(intensity_map, dtype=dtype))
else:
self.k_map = None
@@ -205,7 +205,7 @@ class Bragg2DPtycho(CDIModel):
# The propagation direction of the probe
self.register_buffer('prop_dir',
t.tensor([0, 0, 1], dtype=dtype))
t.as_tensor([0, 0, 1], dtype=dtype))
# This propagator should be able to be multiplied by the propagation
# distance each time to get a propagator
+24 -22
View File
@@ -44,39 +44,39 @@ class FancyPtycho(CDIModel):
super(FancyPtycho, self).__init__()
self.register_buffer('wavelength',
t.tensor(wavelength, dtype=dtype))
t.as_tensor(wavelength, dtype=dtype))
self.store_detector_geometry(detector_geometry,
dtype=dtype)
self.register_buffer('min_translation',
t.tensor(min_translation, dtype=dtype))
t.as_tensor(min_translation, dtype=dtype))
self.register_buffer('obj_basis',
t.tensor(obj_basis, dtype=dtype))
t.as_tensor(obj_basis, dtype=dtype))
if probe_basis is None:
self.register_buffer('probe_basis',
t.tensor(obj_basis, dtype=dtype))
t.as_tensor(obj_basis, dtype=dtype))
else:
self.register_buffer('probe_basis',
t.tensor(probe_basis, dtype=dtype))
t.as_tensor(probe_basis, dtype=dtype))
self.register_buffer('surface_normal',
t.tensor(surface_normal, dtype=dtype))
t.as_tensor(surface_normal, dtype=dtype))
if saturation is None:
self.saturation = None
else:
self.register_buffer('saturation',
t.tensor(saturation, dtype=dtype))
t.as_tensor(saturation, dtype=dtype))
self.register_buffer('fourier_probe',
t.tensor(fourier_probe, dtype=bool))
t.as_tensor(fourier_probe, dtype=bool))
self.register_buffer('exponentiate_obj',
t.tensor(exponentiate_obj, dtype=bool))
t.as_tensor(exponentiate_obj, dtype=bool))
self.register_buffer('phase_only',
t.tensor(phase_only, dtype=bool))
t.as_tensor(phase_only, dtype=bool))
# Not sure how to make this a buffer...
self.units = units
@@ -85,10 +85,10 @@ class FancyPtycho(CDIModel):
self.mask = None
else:
self.register_buffer('mask',
t.tensor(mask, dtype=t.bool))
t.as_tensor(mask, dtype=t.bool))
probe_guess = t.tensor(probe_guess, dtype=t.complex64)
obj_guess = t.tensor(obj_guess, dtype=t.complex64)
probe_guess = t.as_tensor(probe_guess, dtype=t.complex64)
obj_guess = t.as_tensor(obj_guess, dtype=t.complex64)
# We rescale the probe here so it learns at the same rate as the
@@ -121,34 +121,36 @@ class FancyPtycho(CDIModel):
# weights and complex-valued per-mode weight matrices
if len(weights.shape) == 1:
# This is if it's just a list of numbers
self.weights = t.nn.Parameter(t.tensor(weights,
self.weights = t.nn.Parameter(t.as_tensor(weights,
dtype=t.float32))
else:
# Now this is a matrix of weights, so it needs to be complex
self.weights = t.nn.Parameter(t.tensor(weights,
self.weights = t.nn.Parameter(t.as_tensor(weights,
dtype=t.complex64))
if translation_offsets is None:
self.translation_offsets = None
else:
t_o = t.tensor(translation_offsets, dtype=t.float32)
t_o = t.as_tensor(translation_offsets, dtype=t.float32)
t_o = t_o / translation_scale
self.translation_offsets = t.nn.Parameter(t_o)
self.register_buffer('translation_scale',
t.tensor(translation_scale, dtype=dtype))
t.as_tensor(translation_scale, dtype=dtype))
if probe_support is None:
probe_support = t.ones_like(self.probe[0], dtype=t.bool)
self.register_buffer('probe_support',
t.tensor(probe_support, dtype=t.bool))
t.as_tensor(probe_support, dtype=t.bool))
self.probe.data *= self.probe_support
self.register_buffer('oversampling',
t.tensor(oversampling, dtype=int))
t.as_tensor(oversampling, dtype=int))
self.register_buffer('simulate_probe_translation',
t.tensor(simulate_probe_translation, dtype=bool))
self.register_buffer(
'simulate_probe_translation',
t.as_tensor(simulate_probe_translation, dtype=bool)
)
if simulate_probe_translation:
Is = t.arange(self.probe.shape[-2], dtype=dtype)
@@ -162,7 +164,7 @@ class FancyPtycho(CDIModel):
self.register_buffer('simulate_finite_pixels',
t.tensor(simulate_finite_pixels, dtype=bool))
t.as_tensor(simulate_finite_pixels, dtype=bool))
# Here we set the appropriate loss function
if (loss.lower().strip() == 'amplitude mse'
+24 -22
View File
@@ -44,40 +44,40 @@ class MultislicePtycho(CDIModel):
super(MultislicePtycho, self).__init__()
self.register_buffer('wavelength',
t.tensor(wavelength, dtype=dtype))
t.as_tensor(wavelength, dtype=dtype))
self.store_detector_geometry(detector_geometry,
dtype=dtype)
self.register_buffer('min_translation',
t.tensor(min_translation, dtype=dtype))
t.as_tensor(min_translation, dtype=dtype))
self.register_buffer('obj_basis',
t.tensor(obj_basis, dtype=dtype))
t.as_tensor(obj_basis, dtype=dtype))
self.register_buffer('exponentiate_obj',
t.tensor(exponentiate_obj, dtype=bool))
t.as_tensor(exponentiate_obj, dtype=bool))
self.register_buffer('interslice_propagator',
t.tensor(interslice_propagator, dtype=t.complex64))
t.as_tensor(interslice_propagator, dtype=t.complex64))
if probe_basis is None:
self.register_buffer('probe_basis',
t.tensor(obj_basis, dtype=dtype))
t.as_tensor(obj_basis, dtype=dtype))
else:
self.register_buffer('probe_basis',
t.tensor(probe_basis, dtype=dtype))
t.as_tensor(probe_basis, dtype=dtype))
self.register_buffer('surface_normal',
t.tensor(surface_normal, dtype=dtype))
t.as_tensor(surface_normal, dtype=dtype))
if saturation is None:
self.saturation = None
else:
self.register_buffer('saturation',
t.tensor(saturation, dtype=dtype))
t.as_tensor(saturation, dtype=dtype))
self.register_buffer('fourier_probe',
t.tensor(fourier_probe, dtype=bool))
t.as_tensor(fourier_probe, dtype=bool))
# Not sure how to make this a buffer...
self.units = units
@@ -86,10 +86,10 @@ class MultislicePtycho(CDIModel):
self.mask = None
else:
self.register_buffer('mask',
t.tensor(mask, dtype=t.bool))
t.as_tensor(mask, dtype=t.bool))
probe_guess = t.tensor(probe_guess, dtype=t.complex64)
obj_guess = t.tensor(obj_guess, dtype=t.complex64)
probe_guess = t.as_tensor(probe_guess, dtype=t.complex64)
obj_guess = t.as_tensor(obj_guess, dtype=t.complex64)
# We rescale the probe here so it learns at the same rate as the
# object
@@ -121,34 +121,36 @@ class MultislicePtycho(CDIModel):
# weights and complex-valued per-mode weight matrices
if len(weights.shape) == 1:
# This is if it's just a list of numbers
self.weights = t.nn.Parameter(t.tensor(weights,
self.weights = t.nn.Parameter(t.as_tensor(weights,
dtype=t.float32))
else:
# Now this is a matrix of weights, so it needs to be complex
self.weights = t.nn.Parameter(t.tensor(weights,
self.weights = t.nn.Parameter(t.as_tensor(weights,
dtype=t.complex64))
if translation_offsets is None:
self.translation_offsets = None
else:
t_o = t.tensor(translation_offsets, dtype=t.float32)
t_o = t.as_tensor(translation_offsets, dtype=t.float32)
t_o = t_o / translation_scale
self.translation_offsets = t.nn.Parameter(t_o)
self.register_buffer('translation_scale',
t.tensor(translation_scale, dtype=dtype))
t.as_tensor(translation_scale, dtype=dtype))
if probe_support is None:
probe_support = t.ones_like(self.probe[0], dtype=t.bool)
self.register_buffer('probe_support',
t.tensor(probe_support, dtype=t.bool))
t.as_tensor(probe_support, dtype=t.bool))
self.probe.data *= self.probe_support
self.register_buffer('oversampling',
t.tensor(oversampling, dtype=int))
t.as_tensor(oversampling, dtype=int))
self.register_buffer('simulate_probe_translation',
t.tensor(simulate_probe_translation, dtype=bool))
self.register_buffer(
'simulate_probe_translation',
t.as_tensor(simulate_probe_translation, dtype=bool)
)
if simulate_probe_translation:
Is = t.arange(self.probe.shape[-2], dtype=dtype)
@@ -161,7 +163,7 @@ class MultislicePtycho(CDIModel):
self.register_buffer('J_phase', J_phase)
self.register_buffer('simulate_finite_pixels',
t.tensor(simulate_finite_pixels, dtype=bool))
t.as_tensor(simulate_finite_pixels, dtype=bool))
# Here we set the appropriate loss function
if (loss.lower().strip() == 'amplitude mse'
+15 -29
View File
@@ -66,7 +66,7 @@ class RPI(CDIModel):
1j * t.ones([1], dtype=dtype)).dtype
self.register_buffer('wavelength',
t.tensor(wavelength, dtype=dtype))
t.as_tensor(wavelength, dtype=dtype))
self.store_detector_geometry(detector_geometry,
dtype=dtype)
@@ -75,9 +75,9 @@ class RPI(CDIModel):
# used a bandlimiting constraint and had a larger basis, the user is
# expected to upsample it explicitly before doing RPI.
self.register_buffer('probe_basis',
t.tensor(probe_basis, dtype=dtype))
t.as_tensor(probe_basis, dtype=dtype))
scale_factor = t.tensor([probe.shape[-1]/obj_guess.shape[-1],
scale_factor = t.as_tensor([probe.shape[-1]/obj_guess.shape[-1],
probe.shape[-2]/obj_guess.shape[-2]])
self.register_buffer('obj_basis',
(self.probe_basis * scale_factor).to(dtype=dtype))
@@ -86,7 +86,7 @@ class RPI(CDIModel):
self.saturation = None
else:
self.register_buffer('saturation',
t.tensor(saturation, dtype=dtype))
t.as_tensor(saturation, dtype=dtype))
# not sure how to make this a buffer, or if I have to
self.units = units
@@ -95,23 +95,23 @@ class RPI(CDIModel):
self.mask = None
else:
self.register_buffer('mask',
t.tensor(mask, dtype=t.bool))
t.as_tensor(mask, dtype=t.bool))
self.register_buffer('probe', t.tensor(probe, dtype=complex_dtype))
self.register_buffer('probe', t.as_tensor(probe, dtype=complex_dtype))
self.register_buffer('exponentiate_obj',
t.tensor(exponentiate_obj, dtype=bool))
t.as_tensor(exponentiate_obj, dtype=bool))
self.register_buffer('phase_only',
t.tensor(phase_only, dtype=bool))
t.as_tensor(phase_only, dtype=bool))
# We always use multi-modes to store the object, so we convert it
# if we just get a single 2D array as an input
if obj_guess.dim() == 2:
obj_guess = obj_guess[None, :, :]
self.obj = t.nn.Parameter(t.tensor(obj_guess, dtype=complex_dtype))
self.obj = t.nn.Parameter(t.as_tensor(obj_guess, dtype=complex_dtype))
self.weights = t.nn.Parameter(
t.eye(probe.shape[0], dtype=complex_dtype))
@@ -124,40 +124,26 @@ class RPI(CDIModel):
dtype=dtype)
self.register_buffer('background',
t.tensor(background, dtype=t.float32))
t.as_tensor(background, dtype=t.float32))
if obj_support is None:
obj_support = t.ones_like(self.obj[0, ...], dtype=int)
self.register_buffer('obj_support',
t.tensor(obj_support, dtype=int))
t.as_tensor(obj_support, dtype=int))
self.obj.data = self.obj * self.obj_support[None, ...]
self.register_buffer('oversampling',
t.tensor(oversampling, dtype=int))
t.as_tensor(oversampling, dtype=int))
self.register_buffer('propagation_distance',
t.tensor(propagation_distance, dtype=dtype))
t.as_tensor(propagation_distance, dtype=dtype))
# The propagation direction of the probe. For now it's fixed,
# but perhaps it would need to be updated in the future
self.register_buffer('prop_dir',
t.tensor([0, 0, 1], dtype=dtype))
# This propagator should be able to be multiplied by the propagation
# distance each time to get a propagator
#universal_propagator = t.angle(ggasp(
# self.probe.shape[-2:],
# self.probe_basis, self.wavelength,
# t.tensor([0, 0, self.wavelength/(2*np.pi)], dtype=dtype),
# propagation_vector=self.prop_dir,
# dtype=complex_dtype,
# propagate_along_offset=True))
# TODO: probably doesn't support non-float-32 dtypes
#self.register_buffer('universal_propagator',
# universal_propagator)
t.as_tensor([0, 0, 1], dtype=dtype))
@classmethod
@@ -229,7 +215,7 @@ class RPI(CDIModel):
and dataset.background is not None:
background = t.sqrt(dataset.background)
elif background is not None:
background = t.sqrt(t.Tensor(background).to(dtype=t.float32))
background = t.sqrt(t.as_tensor(background).to(dtype=t.float32))
det_geo = dataset.detector_geometry
+2 -5
View File
@@ -13,9 +13,6 @@ __all__ = ['SimplePtycho']
class SimplePtycho(CDIModel):
"""A simple ptychography model for exploring ideas and extensions
"""
def __init__(
self,
@@ -40,8 +37,8 @@ class SimplePtycho(CDIModel):
self.register_buffer('min_translation', t.as_tensor(min_translation))
self.register_buffer('probe_basis', t.as_tensor(probe_basis))
probe_guess = t.tensor(probe_guess, dtype=t.complex64)
obj_guess = t.tensor(obj_guess, dtype=t.complex64)
probe_guess = t.as_tensor(probe_guess, dtype=t.complex64)
obj_guess = t.as_tensor(obj_guess, dtype=t.complex64)
# We rescale the probe here so it learns at the same rate as the
# object
@@ -464,8 +464,8 @@ def ptycho_2D_sinc(probe, obj, translations, shift_probe=True, padding=10, multi
shifted_probe = t.fft.ifft2(t.fft.ifftshift(shifted_fft_probe,
dim=(-1,-2)))
if probe_support is not None:
shifted_probe = shifted_probe * probe_support[..., :, :]
# Note: resist the temptation to remultiply by the probe support here,
# it will fail if you have a probe which is restricted in Fourier space
# TODO This is a kludge, I will fix this.
if multiple_modes and len(selections.shape) == 3: # Multi-mode probe
+48 -1
View File
@@ -14,8 +14,41 @@ import datetime
def pytest_addoption(parser):
parser.addoption(
"--plot", action="store", default=False, help="plot: True to show test plots"
"--plot",
action="store_true",
default=False,
help="when set, shows the test plots"
)
parser.addoption(
"--reconstruction_device",
action="store",
default="cuda",
help="What device to run reconstructions on, if they are being run"
)
parser.addoption(
"--runslow",
action="store_true",
default=False,
help="run slow tests, primarily full reconstruction tests."
)
def pytest_configure(config):
config.addinivalue_line("markers", "slow: mark test as slow to run")
def pytest_collection_modifyitems(config, items):
if config.getoption("--runslow"):
# --runslow given in cli: do not skip slow tests
return
skip_slow = pytest.mark.skip(reason="need --runslow option to run")
for item in items:
if "slow" in item.keywords:
item.add_marker(skip_slow)
@pytest.fixture
def reconstruction_device(request):
return request.config.getoption("--reconstruction_device")
@pytest.fixture
@@ -319,3 +352,17 @@ def test_ptycho_cxis(ptycho_cxi_1, ptycho_cxi_2, ptycho_cxi_3):
on the cxi files.
"""
return [ptycho_cxi_1, ptycho_cxi_2, ptycho_cxi_3]
@pytest.fixture(scope='module')
def gold_ball_cxi(pytestconfig):
return str(pytestconfig.rootpath) + \
'/examples/example_data/AuBalls_700ms_30nmStep_3_6SS_filter.cxi'
@pytest.fixture(scope='module')
def lab_ptycho_cxi(pytestconfig):
return str(pytestconfig.rootpath) + \
'/examples/example_data/lab_ptycho_data.cxi'
+108
View File
@@ -0,0 +1,108 @@
import pytest
import cdtools
import torch as t
import cdtools
from matplotlib import pyplot as plt
@pytest.mark.slow
def test_lab_ptycho(lab_ptycho_cxi, reconstruction_device, show_plot):
print('\nTesting performance on the standard transmission ptycho dataset')
dataset = cdtools.datasets.Ptycho2DDataset.from_cxi(lab_ptycho_cxi)
model = cdtools.models.FancyPtycho.from_dataset(
dataset,
n_modes=3,
oversampling=2,
exponentiate_obj=True,
dm_rank=2,
probe_support_radius=120,
propagation_distance=5e-3,
units='mm',
obj_view_crop=-50,
)
print('Running reconstruction on provided reconstruction_device,',
reconstruction_device)
model.to(device=reconstruction_device)
dataset.get_as(device=reconstruction_device)
for loss in model.Adam_optimize(50, dataset, lr=0.02, batch_size=10):
print(model.report())
if show_plot and model.epoch % 10 == 0:
model.inspect(dataset)
for loss in model.Adam_optimize(50, dataset, lr=0.005, batch_size=50):
print(model.report())
if show_plot and model.epoch % 10 == 0:
model.inspect(dataset)
model.tidy_probes()
if show_plot:
model.inspect(dataset)
model.compare(dataset)
# If this fails, the reconstruction has gotten worse
assert model.loss_history[-1] < 0.001
@pytest.mark.slow
def test_gold_balls(gold_ball_cxi, reconstruction_device, show_plot):
print('\nTesting performance on the standard gold balls dataset')
dataset = cdtools.datasets.Ptycho2DDataset.from_cxi(gold_ball_cxi)
pad = 10
dataset.pad(pad)
model = cdtools.models.FancyPtycho.from_dataset(
dataset,
n_modes=3,
probe_support_radius=50,
propagation_distance=2e-6,
units='um',
probe_fourier_crop=pad
)
model.translation_offsets.data += \
0.7 * t.randn_like(model.translation_offsets)
# Not much probe intensity instability in this dataset, no need for this
model.weights.requires_grad = False
print('Running reconstruction on provided --reconstruction_device,',
reconstruction_device)
model.to(device=reconstruction_device)
dataset.get_as(device=reconstruction_device)
for loss in model.Adam_optimize(20, dataset, lr=0.005, batch_size=50):
print(model.report())
if show_plot and model.epoch % 10 == 0:
model.inspect(dataset)
for loss in model.Adam_optimize(50, dataset, lr=0.002, batch_size=100):
print(model.report())
if show_plot and model.epoch % 10 == 0:
model.inspect(dataset)
for loss in model.Adam_optimize(100, dataset, lr=0.001, batch_size=100,
schedule=True):
print(model.report())
if show_plot and model.epoch % 10 == 0:
model.inspect(dataset)
model.tidy_probes()
if show_plot:
model.inspect(dataset)
model.compare(dataset)
# This just comes from running a reconstruction when it was working well
# and choosing a rough value. If it triggers this assertion error,
# something changed to make the final quality worse!
assert model.loss_history[-1] < 0.0001
+22
View File
@@ -0,0 +1,22 @@
import cdtools
from matplotlib import pyplot as plt
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)
for loss in model.Adam_optimize(100, dataset, batch_size=10):
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