First pass of high NA ptycho and RPI

This commit is contained in:
2025-07-21 11:26:59 +02:00
parent fc5a76ad61
commit 22a28f4043
2 changed files with 71 additions and 6 deletions
+36 -3
View File
@@ -40,6 +40,7 @@ class FancyPtycho(CDIModel):
simulate_finite_pixels=False,
exponentiate_obj=False,
phase_only=False,
high_NA=False,
dtype=t.float32,
obj_view_crop=0
):
@@ -79,6 +80,8 @@ class FancyPtycho(CDIModel):
self.register_buffer('phase_only',
t.as_tensor(phase_only, dtype=bool))
self.register_buffer('high_NA',
t.as_tensor(high_NA, dtype=bool))
# Not sure how to make this a buffer...
self.units = units
@@ -134,6 +137,27 @@ class FancyPtycho(CDIModel):
self.background = t.nn.Parameter(background)
if high_NA:
k_map, intensity_map = \
tools.propagators.generate_high_NA_k_intensity_map(
self.obj_basis,
self.get_detector_geometry()['basis'] / oversampling,
[oversampling * d for d in self.background.shape],
self.get_detector_geometry()['distance'],
self.wavelength,
dtype=t.float32,
lens=False)
self.register_buffer('k_map',
t.as_tensor(k_map, dtype=dtype))
self.register_buffer('intensity_map',
t.as_tensor(intensity_map, dtype=dtype))
else:
self.k_map = None
self.intensity_map = None
if weights is None:
self.weights = None
else:
@@ -228,6 +252,7 @@ class FancyPtycho(CDIModel):
simulate_finite_pixels=False,
exponentiate_obj=False,
phase_only=False,
high_NA=False,
obj_view_crop=None,
obj_padding=200,
):
@@ -288,7 +313,7 @@ class FancyPtycho(CDIModel):
pix_translations,
padding=obj_padding,
)
# Finally, initialize the probe and object using this information
if probe_shape is None:
probe = tools.initializers.SHARP_style_probe(
@@ -436,6 +461,7 @@ class FancyPtycho(CDIModel):
simulate_finite_pixels=simulate_finite_pixels,
phase_only=phase_only,
exponentiate_obj=exponentiate_obj,
high_NA=high_NA,
obj_view_crop=obj_view_crop
)
@@ -540,11 +566,18 @@ class FancyPtycho(CDIModel):
def forward_propagator(self, wavefields):
return tools.propagators.far_field(wavefields)
if self.high_NA:
return tools.propagators.high_NA_far_field(
wavefields,self.k_map,intensity_map=self.intensity_map)
else:
return tools.propagators.far_field(wavefields)
def backward_propagator(self, wavefields):
return tools.propagators.inverse_far_field(wavefields)
if self.high_NA:
assert NotImplementedError('Backward propagator not defined with tilt correction')
else:
return tools.propagators.inverse_far_field(wavefields)
def measurement(self, wavefields):
+35 -3
View File
@@ -55,6 +55,7 @@ class RPI(CDIModel):
weight_matrix=False,
exponentiate_obj=False,
phase_only=False,
high_NA=False,
propagation_distance=0,
units='um',
dtype=t.float32,
@@ -106,13 +107,17 @@ class RPI(CDIModel):
self.register_buffer('phase_only',
t.as_tensor(phase_only, dtype=bool))
self.register_buffer('high_NA',
t.as_tensor(high_NA, 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.as_tensor(obj_guess, dtype=complex_dtype))
self.weights = t.nn.Parameter(
t.eye(probe.shape[0], dtype=complex_dtype))
@@ -145,6 +150,27 @@ class RPI(CDIModel):
self.register_buffer('prop_dir',
t.as_tensor([0, 0, 1], dtype=dtype))
if high_NA:
k_map, intensity_map = \
tools.propagators.generate_high_NA_k_intensity_map(
self.probe_basis,
self.get_detector_geometry()['basis'] / oversampling,
[oversampling * d for d in self.background.shape],
self.get_detector_geometry()['distance'],
self.wavelength,
dtype=t.float32,
lens=False)
self.register_buffer('k_map',
t.as_tensor(k_map, dtype=dtype))
self.register_buffer('intensity_map',
t.as_tensor(intensity_map, dtype=dtype))
else:
self.k_map = None
self.intensity_map = None
@classmethod
def from_dataset(
@@ -162,6 +188,7 @@ class RPI(CDIModel):
weight_matrix=False,
exponentiate_obj=False,
phase_only=False,
high_NA=False,
probe_threshold=0,
dtype=t.float32,
):
@@ -254,6 +281,7 @@ class RPI(CDIModel):
oversampling=oversampling,
exponentiate_obj=exponentiate_obj,
phase_only=phase_only,
high_NA=high_NA,
weight_matrix=weight_matrix)
# I don't love this pattern, where I do the "real" obj initialization
@@ -497,8 +525,12 @@ class RPI(CDIModel):
def forward_propagator(self, wavefields):
p = tools.propagators.far_field(wavefields)
return p
if self.high_NA:
wavefields = wavefields.reshape([wavefields.shape[0]] + [-1] + list(wavefields.shape[-2:]))
return tools.propagators.high_NA_far_field(
wavefields,self.k_map,intensity_map=self.intensity_map)
else:
return tools.propagators.far_field(wavefields)
def backward_propagator(self, wavefields):