mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-09-09 13:02:41 +02:00
Merge pull request #71 from cdtools-developers/bugfix/intensity_loading
Fix issue #68 for loading of datasets with intensities when using OPRP
This commit is contained in:
@@ -459,7 +459,8 @@ class FancyPtycho(CDIModel):
|
||||
if hasattr(dataset, 'intensities') and dataset.intensities is not None:
|
||||
intensities = dataset.intensities.to(dtype=Ws.dtype)[:,...]
|
||||
weights = t.sqrt(intensities)
|
||||
Ws *= (weights / t.mean(weights))
|
||||
Ws *= (weights / t.mean(weights)).reshape(
|
||||
(len(weights),) + (1,)*(Ws.ndim - 1))
|
||||
|
||||
if hasattr(dataset, 'mask') and dataset.mask is not None:
|
||||
mask = dataset.mask.to(t.bool)
|
||||
|
||||
@@ -45,6 +45,28 @@ def test_center_probe(lab_ptycho_cxi):
|
||||
rtol=1e-3
|
||||
)
|
||||
|
||||
def test_lab_ptycho_data_loading(lab_ptycho_cxi):
|
||||
|
||||
print('\nTesting a few unusual data loading scenarios.')
|
||||
dataset = cdtools.datasets.Ptycho2DDataset.from_cxi(lab_ptycho_cxi)
|
||||
|
||||
# Test that it will properly load an initialization for the weights
|
||||
# from the intensities with OPRP on
|
||||
dataset.intensities = t.rand(len(dataset))
|
||||
|
||||
model = cdtools.models.FancyPtycho.from_dataset(
|
||||
dataset,
|
||||
n_modes=4,
|
||||
dm_rank=1,
|
||||
)
|
||||
|
||||
# And test the case without OPRP
|
||||
model = cdtools.models.FancyPtycho.from_dataset(
|
||||
dataset,
|
||||
n_modes=2,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@pytest.mark.slow
|
||||
def test_lab_ptycho(lab_ptycho_cxi, reconstruction_device, show_plot):
|
||||
@@ -71,8 +93,8 @@ def test_lab_ptycho(lab_ptycho_cxi, reconstruction_device, show_plot):
|
||||
|
||||
print('Running reconstruction on provided reconstruction_device,',
|
||||
reconstruction_device)
|
||||
model.to(device=reconstruction_device)
|
||||
dataset.get_as(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())
|
||||
|
||||
Reference in New Issue
Block a user