diff --git a/CDTools/models/simple_ptycho.py b/CDTools/models/simple_ptycho.py index b0ac9ba..dc9d27f 100644 --- a/CDTools/models/simple_ptycho.py +++ b/CDTools/models/simple_ptycho.py @@ -4,6 +4,7 @@ import torch as t from CDTools.models import CDIModel from CDTools import tools from copy import copy +from torch.utils import data as torchdata class SimplePtycho(CDIModel): @@ -134,7 +135,7 @@ class SimplePtycho(CDIModel): pass - def ePIE(iterations, dataset, beta = 1.0): + def ePIE(self, iterations, dataset, beta = 1.0): """Runs an ePIE reconstruction as described in `Maiden et al. (2017) `_. Optional parameters are: @@ -144,21 +145,34 @@ class SimplePtycho(CDIModel): :arg ``object``: Initial object wavefunction. """ probe_shape = self.probe.shape + + if self.mask is not None: + mask = self.mask[...,None] + def probe_update(exit_wave, exit_wave_corrected, probe, object, translation): - return probe+tools.cmult(beta*(tools.cconj(object)/ \ - t.max(tools.cabssq(object)))[translations[0]:translations[0]+probe_shape[0],translations[1]:translations[1]+probe_shape[1]], \ + return probe+tools.cmath.cmult(beta*(tools.cmath.cconj(object)/t.max(tools.cmath.cabssq(object)))[translation[0]:translation[0]+probe_shape[0],translation[1]:translation[1]+probe_shape[1]], \ exit_wave_corrected-exit_wave) + def object_update(exit_wave, exit_wave_corrected, probe, object, translation): - object[translations[0]:translations[0]+probe_shape[0],translations[1]:translations[1]+probe_shape[1]]\ - +=tools.cmult(tools.cconj(probe)/t.max(tools.cabssq(probe)),exit_wave_corrected-exit_wave) + object[translation[0]:translation[0]+probe_shape[0],translation[1]:translation[1]+probe_shape[1]]\ + +=tools.cmath.cmult(tools.cmath.cconj(probe)/t.max(tools.cmath.cabssq(probe)),exit_wave_corrected-exit_wave) return object + with t.no_grad(): data_loader = torchdata.DataLoader(dataset, shuffle=True) + for it in range(iterations): - for translations, patterns in data_loader: + loss = [] + for (i, [translations]), [patterns] in data_loader: probe = self.probe.clone() object = self.obj.clone() - exit_wave = tools.cmult(probe, object[translations:translations[0]+probe_shape[0],translations[1]:translations[1]+probe_shape[1]]).clone() - exit_wave_corrected = projectors.modulus(self.forward_propagator(exit_wave), patterns, mask = self.mask) - self.probe = probe_update(exit_wave, exit_wave_corrected, probe, object, translations) - self.obj = object_update(exit_wave, exit_wave_corrected, probe, object, translations) + exit_wave = self.interaction(i, translations) + exit_wave_corrected = exit_wave.clone() + exit_wave_corrected[self.detector_slice] = tools.projectors.modulus(self.forward_propagator(exit_wave)[self.detector_slice], patterns, mask = mask) + + integer_translations = t.round(translations).to(dtype=t.int32) + self.probe.data = probe_update(exit_wave, exit_wave_corrected, probe, object, integer_translations) + self.obj.data = object_update(exit_wave, exit_wave_corrected, probe, object, integer_translations) + loss.append(self.loss(self.measurement(self.interaction(i, translations)), patterns)) + + yield t.mean(t.Tensor(loss)).cpu().numpy() diff --git a/examples/simple_ptycho.py b/examples/simple_ptycho.py index 72ef0c6..f5a0f19 100644 --- a/examples/simple_ptycho.py +++ b/examples/simple_ptycho.py @@ -7,9 +7,11 @@ import h5py import torch as t import numpy as np -filename = '../../../Projects/CSX_3_19/cxis/processed/114429_p.cxi' -filename = '../../../Projects/CSX_3_19/cxis/processed/115145_p.cxi' -filename = '../../../Downloads/AuBalls_700ms_30nmStep_3_3SS_filter.cxi' + +filename = '../../Downloads/114429_p.cxi' +#filename = '../../../Projects/CSX_3_19/cxis/processed/114429_p.cxi' +#filename = '../../../Projects/CSX_3_19/cxis/processed/115145_p.cxi' +#filename = '../../../Downloads/AuBalls_700ms_30nmStep_3_3SS_filter.cxi' with h5py.File(filename,'r') as f: @@ -21,12 +23,12 @@ model = CDTools.models.SimplePtycho.from_dataset(dataset) # Uncomment these to use on the CPU # default is CPU with 32-bit floats -model.to(device='cuda') +#model.to(device='cuda') #dataset.to(device='cuda') -dataset.get_as(device='cuda') +#dataset.get_as(device='cuda') -for loss in model.ePIE(100, dataset): +for loss in model.ePIE(1, dataset): print(loss) from matplotlib import pyplot as plt