phaserEM 0.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- phaser/__init__.py +0 -0
- phaser/__main__.py +5 -0
- phaser/engines/common/__init__.py +0 -0
- phaser/engines/common/noise_models.py +113 -0
- phaser/engines/common/output.py +189 -0
- phaser/engines/common/position_correction.py +62 -0
- phaser/engines/common/regularizers.py +403 -0
- phaser/engines/common/simulation.py +270 -0
- phaser/engines/conventional/__init__.py +0 -0
- phaser/engines/conventional/run.py +142 -0
- phaser/engines/conventional/solvers.py +476 -0
- phaser/engines/gradient/run.py +451 -0
- phaser/engines/gradient/solvers.py +139 -0
- phaser/execute.py +371 -0
- phaser/hooks/__init__.py +158 -0
- phaser/hooks/hook.py +159 -0
- phaser/hooks/io/empad.py +88 -0
- phaser/hooks/object.py +25 -0
- phaser/hooks/preprocessing.py +133 -0
- phaser/hooks/probe.py +24 -0
- phaser/hooks/regularization.py +97 -0
- phaser/hooks/scan.py +27 -0
- phaser/hooks/schedule.py +75 -0
- phaser/hooks/solver.py +169 -0
- phaser/io/__init__.py +0 -0
- phaser/io/empad.py +212 -0
- phaser/main.py +92 -0
- phaser/plan.py +184 -0
- phaser/py.typed +0 -0
- phaser/state.py +249 -0
- phaser/types.py +305 -0
- phaser/utils/__init__.py +0 -0
- phaser/utils/_cuda_kernels.py +213 -0
- phaser/utils/_jax_kernels.py +98 -0
- phaser/utils/analysis.py +263 -0
- phaser/utils/image.py +201 -0
- phaser/utils/io.py +402 -0
- phaser/utils/misc.py +295 -0
- phaser/utils/num.py +800 -0
- phaser/utils/object.py +578 -0
- phaser/utils/optics.py +377 -0
- phaser/utils/physics.py +88 -0
- phaser/utils/plotting.py +699 -0
- phaser/utils/scan.py +60 -0
- phaser/web/__init__.py +0 -0
- phaser/web/dist/03510a839ccb97b0da9f.module.wasm +0 -0
- phaser/web/dist/9573273f862f4f5d9644.module.wasm +0 -0
- phaser/web/dist/bundle-dashboard.js +712 -0
- phaser/web/dist/bundle-manager.js +210 -0
- phaser/web/dist/bundle-vendors-node_modules_wasm-array_wasm_array_js.js +106 -0
- phaser/web/dist/style.css +152 -0
- phaser/web/notebook.py +269 -0
- phaser/web/routes.py +211 -0
- phaser/web/server.py +540 -0
- phaser/web/slurm.py +180 -0
- phaser/web/templates/base.html +13 -0
- phaser/web/templates/dashboard.html +14 -0
- phaser/web/templates/manager.html +19 -0
- phaser/web/types.py +255 -0
- phaser/web/util.py +116 -0
- phaser/web/worker.py +200 -0
- phaserem-0.1.dist-info/METADATA +121 -0
- phaserem-0.1.dist-info/RECORD +67 -0
- phaserem-0.1.dist-info/WHEEL +5 -0
- phaserem-0.1.dist-info/entry_points.txt +2 -0
- phaserem-0.1.dist-info/licenses/LICENSE.txt +373 -0
- phaserem-0.1.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,403 @@
|
|
|
1
|
+
from functools import partial
|
|
2
|
+
import logging
|
|
3
|
+
import typing as t
|
|
4
|
+
|
|
5
|
+
import numpy
|
|
6
|
+
from numpy.typing import NDArray
|
|
7
|
+
|
|
8
|
+
from phaser.utils.num import (
|
|
9
|
+
get_array_module, get_scipy_module, Float,
|
|
10
|
+
jit, fft2, ifft2, abs2, xp_is_jax, to_real_dtype
|
|
11
|
+
)
|
|
12
|
+
from phaser.state import ReconsState
|
|
13
|
+
from phaser.hooks.regularization import (
|
|
14
|
+
ClampObjectAmplitudeProps, LimitProbeSupportProps,
|
|
15
|
+
RegularizeLayersProps, ObjLowPassProps,
|
|
16
|
+
CostRegularizerProps, TVRegularizerProps
|
|
17
|
+
)
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class ClampObjectAmplitude:
|
|
21
|
+
def __init__(self, args: None, props: ClampObjectAmplitudeProps):
|
|
22
|
+
self.amplitude = props.amplitude
|
|
23
|
+
|
|
24
|
+
def init_state(self, sim: ReconsState) -> None:
|
|
25
|
+
return None
|
|
26
|
+
|
|
27
|
+
def apply_group(self, group: NDArray[numpy.integer], sim: ReconsState, state: None) -> t.Tuple[ReconsState, None]:
|
|
28
|
+
return self.apply_iter(sim, state)
|
|
29
|
+
|
|
30
|
+
def apply_iter(self, sim: ReconsState, state: None) -> t.Tuple[ReconsState, None]:
|
|
31
|
+
amp = to_real_dtype(sim.object.data.dtype)(self.amplitude)
|
|
32
|
+
sim.object.data = clamp_amplitude(sim.object.data, amp)
|
|
33
|
+
return (sim, None)
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
@partial(jit, donate_argnames=('obj',), cupy_fuse=True)
|
|
37
|
+
def clamp_amplitude(obj: NDArray[numpy.complexfloating], amplitude: t.Union[float, numpy.floating]) -> NDArray[numpy.complexfloating]:
|
|
38
|
+
xp = get_array_module(obj)
|
|
39
|
+
|
|
40
|
+
obj_amp = xp.abs(obj)
|
|
41
|
+
scale = xp.minimum(obj_amp, amplitude) / obj_amp
|
|
42
|
+
return obj * scale
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
class LimitProbeSupport:
|
|
46
|
+
def __init__(self, args: None, props: LimitProbeSupportProps):
|
|
47
|
+
self.max_angle = props.max_angle
|
|
48
|
+
|
|
49
|
+
def init_state(self, sim: ReconsState) -> NDArray[numpy.bool_]:
|
|
50
|
+
xp = get_array_module(sim.probe.data)
|
|
51
|
+
(ky, kx) = sim.probe.sampling.recip_grid(xp=xp)
|
|
52
|
+
mask = kx**2 + ky**2 <= (self.max_angle*1e-3 / sim.wavelength)**2
|
|
53
|
+
return mask
|
|
54
|
+
|
|
55
|
+
def apply_group(self, group: NDArray[numpy.integer], sim: ReconsState, state: NDArray[numpy.bool_]) -> t.Tuple[ReconsState, NDArray[numpy.bool_]]:
|
|
56
|
+
return self.apply_iter(sim, state)
|
|
57
|
+
|
|
58
|
+
def apply_iter(self, sim: ReconsState, state: NDArray[numpy.bool_]) -> t.Tuple[ReconsState, NDArray[numpy.bool_]]:
|
|
59
|
+
mask = state
|
|
60
|
+
#xp = get_array_module(sim.state.probe.data)
|
|
61
|
+
#print(f"intensity before: {xp.sum(abs2(sim.state.probe.data))}")
|
|
62
|
+
sim.probe.data = ifft2(fft2(sim.probe.data) * mask)
|
|
63
|
+
#print(f"intensity after: {xp.sum(abs2(sim.state.probe.data))}")
|
|
64
|
+
return (sim, mask)
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
class RemovePhaseRamp:
|
|
68
|
+
def __init__(self, args: None, props: t.Any):
|
|
69
|
+
...
|
|
70
|
+
|
|
71
|
+
def init_state(self, sim: ReconsState) -> NDArray[numpy.bool_]:
|
|
72
|
+
xp = get_array_module(sim.object.data)
|
|
73
|
+
return sim.object.sampling.get_region_mask(xp=xp)
|
|
74
|
+
|
|
75
|
+
def apply_group(self, group: NDArray[numpy.integer], sim: ReconsState, state: NDArray[numpy.bool_]) -> t.Tuple[ReconsState, NDArray[numpy.bool_]]:
|
|
76
|
+
return self.apply_iter(sim, state)
|
|
77
|
+
|
|
78
|
+
def apply_iter(self, sim: ReconsState, state: NDArray[numpy.bool_]) -> t.Tuple[ReconsState, NDArray[numpy.bool_]]:
|
|
79
|
+
from phaser.utils.image import remove_linear_ramp
|
|
80
|
+
xp = get_array_module(sim.object.data)
|
|
81
|
+
phase = remove_linear_ramp(xp.angle(sim.object.data), state)
|
|
82
|
+
sim.object.data = t.cast(NDArray[numpy.complexfloating], xp.abs(sim.object.data) * xp.exp(1.j * phase))
|
|
83
|
+
return (sim, state)
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
class RegularizeLayers:
|
|
87
|
+
def __init__(self, args: None, props: RegularizeLayersProps):
|
|
88
|
+
self.weight = props.weight
|
|
89
|
+
self.sigma = props.sigma
|
|
90
|
+
|
|
91
|
+
def init_state(self, sim: ReconsState) -> None:
|
|
92
|
+
return None
|
|
93
|
+
|
|
94
|
+
def apply_iter(self, sim: ReconsState, state: None) -> t.Tuple[ReconsState, None]:
|
|
95
|
+
xp = get_array_module(sim.object.data)
|
|
96
|
+
scipy = get_scipy_module(sim.object.data)
|
|
97
|
+
dtype = to_real_dtype(sim.object.data)
|
|
98
|
+
|
|
99
|
+
if len(sim.object.thicknesses) < 2:
|
|
100
|
+
return (sim, None)
|
|
101
|
+
|
|
102
|
+
# approximate layers as equally spaced
|
|
103
|
+
layer_spacing = numpy.mean(sim.object.thicknesses)
|
|
104
|
+
# calculate size of filter (go to ~sigma in each direction)
|
|
105
|
+
r = int(numpy.ceil(2. * self.sigma / layer_spacing))
|
|
106
|
+
n = 2*r + 1
|
|
107
|
+
|
|
108
|
+
# make Gaussian filter
|
|
109
|
+
zs = ((xp.arange(0, n) - (n-1)//2) * layer_spacing).astype(dtype)
|
|
110
|
+
kernel = xp.exp(-(zs / self.sigma)**2 / 2.)
|
|
111
|
+
kernel /= xp.sum(kernel)
|
|
112
|
+
|
|
113
|
+
# we convolve the log of object, because the transmission
|
|
114
|
+
# function is multiplicative, not additive
|
|
115
|
+
|
|
116
|
+
if xp_is_jax(xp):
|
|
117
|
+
new_obj = xp.exp(scipy.signal.convolve(
|
|
118
|
+
xp.pad(xp.log(sim.object.data), ((r, r), (0, 0), (0, 0)), mode='edge'),
|
|
119
|
+
kernel[:, None, None],
|
|
120
|
+
mode="valid"
|
|
121
|
+
))
|
|
122
|
+
else:
|
|
123
|
+
new_obj = xp.exp(scipy.ndimage.convolve1d(xp.log(
|
|
124
|
+
sim.object.data
|
|
125
|
+
), kernel, axis=0, mode='nearest'))
|
|
126
|
+
|
|
127
|
+
assert new_obj.shape == sim.object.data.shape
|
|
128
|
+
assert new_obj.dtype == sim.object.data.dtype
|
|
129
|
+
sim.object.data = (
|
|
130
|
+
self.weight * new_obj + (1 - self.weight) * sim.object.data
|
|
131
|
+
)
|
|
132
|
+
return (sim, None)
|
|
133
|
+
|
|
134
|
+
|
|
135
|
+
class ObjLowPass:
|
|
136
|
+
def __init__(self, args: None, props: ObjLowPassProps):
|
|
137
|
+
self.logger = logging.getLogger(__name__)
|
|
138
|
+
self.max_freq = props.max_freq
|
|
139
|
+
|
|
140
|
+
def init_state(self, sim: ReconsState) -> NDArray[numpy.bool_]:
|
|
141
|
+
samp = sim.object.sampling
|
|
142
|
+
xp = get_array_module(sim.object.data)
|
|
143
|
+
|
|
144
|
+
ky = xp.fft.fftfreq(samp.shape[0], 1.0)
|
|
145
|
+
kx = xp.fft.fftfreq(samp.shape[1], 1.0)
|
|
146
|
+
(ky, kx) = xp.meshgrid(ky, kx, indexing='ij')
|
|
147
|
+
k2 = ky**2 + kx**2
|
|
148
|
+
|
|
149
|
+
return k2 <= self.max_freq**2
|
|
150
|
+
|
|
151
|
+
def apply_group(
|
|
152
|
+
self, group: NDArray[numpy.integer], sim: ReconsState, state: NDArray[numpy.bool_]
|
|
153
|
+
) -> t.Tuple[ReconsState, NDArray[numpy.bool_]]:
|
|
154
|
+
return self.apply_iter(sim, state)
|
|
155
|
+
|
|
156
|
+
def apply_iter(
|
|
157
|
+
self, sim: ReconsState, state: NDArray[numpy.bool_]
|
|
158
|
+
) -> t.Tuple[ReconsState, NDArray[numpy.bool_]]:
|
|
159
|
+
# TODO: should this be done in-place?
|
|
160
|
+
sim.object.data = ifft2(state * fft2(sim.object.data))
|
|
161
|
+
return (sim, state)
|
|
162
|
+
|
|
163
|
+
|
|
164
|
+
class ObjL1:
|
|
165
|
+
def __init__(self, args: None, props: CostRegularizerProps):
|
|
166
|
+
self.cost: float = props.cost
|
|
167
|
+
|
|
168
|
+
def init_state(self, sim: ReconsState) -> None:
|
|
169
|
+
return None
|
|
170
|
+
|
|
171
|
+
def calc_loss_group(
|
|
172
|
+
self, group: NDArray[numpy.integer], sim: ReconsState, state: None
|
|
173
|
+
) -> t.Tuple[Float, None]:
|
|
174
|
+
xp = get_array_module(sim.object.data)
|
|
175
|
+
|
|
176
|
+
cost = xp.sum(xp.abs(sim.object.data - 1.0))
|
|
177
|
+
cost_scale = (group.shape[-1] / numpy.prod(sim.scan.shape[:-1])).astype(cost.dtype)
|
|
178
|
+
return (cost * cost_scale * self.cost, state)
|
|
179
|
+
|
|
180
|
+
|
|
181
|
+
class ObjL2:
|
|
182
|
+
def __init__(self, args: None, props: CostRegularizerProps):
|
|
183
|
+
self.cost: Float = props.cost
|
|
184
|
+
|
|
185
|
+
def init_state(self, sim: ReconsState) -> None:
|
|
186
|
+
return None
|
|
187
|
+
|
|
188
|
+
def calc_loss_group(
|
|
189
|
+
self, group: NDArray[numpy.integer], sim: ReconsState, state: None
|
|
190
|
+
) -> t.Tuple[Float, None]:
|
|
191
|
+
xp = get_array_module(sim.object.data)
|
|
192
|
+
|
|
193
|
+
cost = xp.sum(abs2(sim.object.data - 1.0))
|
|
194
|
+
cost_scale = (group.shape[-1] / numpy.prod(sim.scan.shape[:-1])).astype(cost.dtype)
|
|
195
|
+
return (cost * cost_scale * self.cost, state)
|
|
196
|
+
|
|
197
|
+
|
|
198
|
+
class ObjPhaseL1:
|
|
199
|
+
def __init__(self, args: None, props: CostRegularizerProps):
|
|
200
|
+
self.cost: float = props.cost
|
|
201
|
+
|
|
202
|
+
def init_state(self, sim: ReconsState) -> None:
|
|
203
|
+
return None
|
|
204
|
+
|
|
205
|
+
def calc_loss_group(
|
|
206
|
+
self, group: NDArray[numpy.integer], sim: ReconsState, state: None
|
|
207
|
+
) -> t.Tuple[Float, None]:
|
|
208
|
+
xp = get_array_module(sim.object.data)
|
|
209
|
+
|
|
210
|
+
cost = xp.sum(xp.abs(xp.angle(sim.object.data)))
|
|
211
|
+
cost_scale = (group.shape[-1] / numpy.prod(sim.scan.shape[:-1])).astype(cost.dtype)
|
|
212
|
+
return (cost * cost_scale * self.cost, state)
|
|
213
|
+
|
|
214
|
+
|
|
215
|
+
class ObjRecipL1:
|
|
216
|
+
def __init__(self, args: None, props: CostRegularizerProps):
|
|
217
|
+
self.cost: float = props.cost
|
|
218
|
+
|
|
219
|
+
def init_state(self, sim: ReconsState) -> None:
|
|
220
|
+
return None
|
|
221
|
+
|
|
222
|
+
def calc_loss_group(
|
|
223
|
+
self, group: NDArray[numpy.integer], sim: ReconsState, state: None
|
|
224
|
+
) -> t.Tuple[Float, None]:
|
|
225
|
+
xp = get_array_module(sim.object.data)
|
|
226
|
+
|
|
227
|
+
# l1 norm of diff. pattern amplitude
|
|
228
|
+
# TODO log object before this?
|
|
229
|
+
cost = xp.sum(
|
|
230
|
+
xp.abs(fft2(xp.prod(sim.object.data, axis=0)))
|
|
231
|
+
)
|
|
232
|
+
# scale cost by fraction of the total reconstruction in the group
|
|
233
|
+
cost_scale = (group.shape[-1] / numpy.prod(sim.scan.shape[:-1])).astype(cost.dtype)
|
|
234
|
+
|
|
235
|
+
return (cost * cost_scale * self.cost, state)
|
|
236
|
+
|
|
237
|
+
|
|
238
|
+
class ObjTotalVariation:
|
|
239
|
+
def __init__(self, args: None, props: TVRegularizerProps):
|
|
240
|
+
self.cost: float = props.cost
|
|
241
|
+
self.eps: float = props.eps
|
|
242
|
+
|
|
243
|
+
def init_state(self, sim: ReconsState) -> None:
|
|
244
|
+
return None
|
|
245
|
+
|
|
246
|
+
def calc_loss_group(
|
|
247
|
+
self, group: NDArray[numpy.integer], sim: ReconsState, state: None
|
|
248
|
+
) -> t.Tuple[Float, None]:
|
|
249
|
+
xp = get_array_module(sim.object.data)
|
|
250
|
+
|
|
251
|
+
# isotropic total variation
|
|
252
|
+
g_y, g_x = img_grad(sim.object.data)
|
|
253
|
+
cost = xp.sum(xp.sqrt(abs2(g_y) + abs2(g_x) + self.eps))
|
|
254
|
+
# anisotropic total variation
|
|
255
|
+
#cost = (
|
|
256
|
+
# xp.sum(xp.abs(xp.diff(sim.object.data, axis=-1))) +
|
|
257
|
+
# xp.sum(xp.abs(xp.diff(sim.object.data, axis=-2)))
|
|
258
|
+
#)
|
|
259
|
+
# scale cost by fraction of the total reconstruction in the group
|
|
260
|
+
# TODO also scale by # of pixels or similar?
|
|
261
|
+
cost_scale = (group.shape[-1] / numpy.prod(sim.scan.shape[:-1])).astype(cost.dtype)
|
|
262
|
+
|
|
263
|
+
return (cost * cost_scale * self.cost, state)
|
|
264
|
+
|
|
265
|
+
|
|
266
|
+
class ObjTikhonov:
|
|
267
|
+
def __init__(self, args: None, props: CostRegularizerProps):
|
|
268
|
+
self.cost: float = props.cost
|
|
269
|
+
|
|
270
|
+
def init_state(self, sim: ReconsState) -> None:
|
|
271
|
+
return None
|
|
272
|
+
|
|
273
|
+
def calc_loss_group(
|
|
274
|
+
self, group: NDArray[numpy.integer], sim: ReconsState, state: None
|
|
275
|
+
) -> t.Tuple[Float, None]:
|
|
276
|
+
xp = get_array_module(sim.object.data)
|
|
277
|
+
|
|
278
|
+
cost = (
|
|
279
|
+
xp.sum(abs2(xp.diff(sim.object.data, axis=-1))) +
|
|
280
|
+
xp.sum(abs2(xp.diff(sim.object.data, axis=-2)))
|
|
281
|
+
)
|
|
282
|
+
# scale cost by fraction of the total reconstruction in the group
|
|
283
|
+
cost_scale = (group.shape[-1] / numpy.prod(sim.scan.shape[:-1])).astype(cost.dtype)
|
|
284
|
+
|
|
285
|
+
return (cost * cost_scale * self.cost, state)
|
|
286
|
+
|
|
287
|
+
|
|
288
|
+
class LayersTotalVariation:
|
|
289
|
+
def __init__(self, args: None, props: CostRegularizerProps):
|
|
290
|
+
self.cost: float = props.cost
|
|
291
|
+
|
|
292
|
+
def init_state(self, sim: ReconsState) -> None:
|
|
293
|
+
return None
|
|
294
|
+
|
|
295
|
+
def calc_loss_group(
|
|
296
|
+
self, group: NDArray[numpy.integer], sim: ReconsState, state: None
|
|
297
|
+
) -> t.Tuple[Float, None]:
|
|
298
|
+
xp = get_array_module(sim.object.data)
|
|
299
|
+
|
|
300
|
+
if sim.object.data.shape[0] < 2:
|
|
301
|
+
return (0.0, state)
|
|
302
|
+
|
|
303
|
+
cost = xp.sum(xp.abs(xp.diff(sim.object.data, axis=0)))
|
|
304
|
+
# scale cost by fraction of the total reconstruction in the group
|
|
305
|
+
cost_scale = (group.shape[-1] / numpy.prod(sim.scan.shape[:-1])).astype(cost.dtype)
|
|
306
|
+
|
|
307
|
+
return (cost * cost_scale * self.cost, state)
|
|
308
|
+
|
|
309
|
+
|
|
310
|
+
class LayersTikhonov:
|
|
311
|
+
def __init__(self, args: None, props: CostRegularizerProps):
|
|
312
|
+
self.cost: float = props.cost
|
|
313
|
+
|
|
314
|
+
def init_state(self, sim: ReconsState) -> None:
|
|
315
|
+
return None
|
|
316
|
+
|
|
317
|
+
def calc_loss_group(
|
|
318
|
+
self, group: NDArray[numpy.integer], sim: ReconsState, state: None
|
|
319
|
+
) -> t.Tuple[Float, None]:
|
|
320
|
+
xp = get_array_module(sim.object.data)
|
|
321
|
+
|
|
322
|
+
if sim.object.data.shape[0] < 2:
|
|
323
|
+
return (0.0, state)
|
|
324
|
+
|
|
325
|
+
cost = xp.sum(abs2(xp.diff(sim.object.data, axis=0)))
|
|
326
|
+
# scale cost by fraction of the total reconstruction in the group
|
|
327
|
+
cost_scale = (group.shape[-1] / numpy.prod(sim.scan.shape[:-1])).astype(cost.dtype)
|
|
328
|
+
|
|
329
|
+
return (cost * cost_scale * self.cost, state)
|
|
330
|
+
|
|
331
|
+
|
|
332
|
+
class ProbePhaseTikhonov:
|
|
333
|
+
def __init__(self, args: None, props: CostRegularizerProps):
|
|
334
|
+
self.cost: float = props.cost
|
|
335
|
+
|
|
336
|
+
def init_state(self, sim: ReconsState) -> None:
|
|
337
|
+
return None
|
|
338
|
+
|
|
339
|
+
def calc_loss_group(
|
|
340
|
+
self, group: NDArray[numpy.integer], sim: ReconsState, state: None
|
|
341
|
+
) -> t.Tuple[Float, None]:
|
|
342
|
+
xp = get_array_module(sim.probe.data)
|
|
343
|
+
|
|
344
|
+
phase = xp.angle(fft2(sim.probe.data))
|
|
345
|
+
|
|
346
|
+
cost = (
|
|
347
|
+
xp.sum(abs2(xp.diff(phase, axis=-1))) +
|
|
348
|
+
xp.sum(abs2(xp.diff(phase, axis=-2)))
|
|
349
|
+
)
|
|
350
|
+
cost_scale = 1.0
|
|
351
|
+
|
|
352
|
+
return (cost * cost_scale * self.cost, state)
|
|
353
|
+
|
|
354
|
+
|
|
355
|
+
class ProbeRecipTikhonov:
|
|
356
|
+
def __init__(self, args: None, props: CostRegularizerProps):
|
|
357
|
+
self.cost: float = props.cost
|
|
358
|
+
|
|
359
|
+
def init_state(self, sim: ReconsState) -> None:
|
|
360
|
+
return None
|
|
361
|
+
|
|
362
|
+
def calc_loss_group(
|
|
363
|
+
self, group: NDArray[numpy.integer], sim: ReconsState, state: None
|
|
364
|
+
) -> t.Tuple[Float, None]:
|
|
365
|
+
xp = get_array_module(sim.probe.data)
|
|
366
|
+
probe_recip = xp.fft.fftshift(fft2(sim.probe.data), axes=(-1, -2))
|
|
367
|
+
|
|
368
|
+
cost = (
|
|
369
|
+
xp.sum(abs2(xp.diff(probe_recip, axis=-1))) +
|
|
370
|
+
xp.sum(abs2(xp.diff(probe_recip, axis=-2)))
|
|
371
|
+
)
|
|
372
|
+
cost_scale = 1.0
|
|
373
|
+
|
|
374
|
+
return (cost * cost_scale * self.cost, state)
|
|
375
|
+
|
|
376
|
+
|
|
377
|
+
class ProbeRecipTotalVariation:
|
|
378
|
+
def __init__(self, args: None, props: TVRegularizerProps):
|
|
379
|
+
self.cost: float = props.cost
|
|
380
|
+
self.eps: float = props.eps
|
|
381
|
+
|
|
382
|
+
def init_state(self, sim: ReconsState) -> None:
|
|
383
|
+
return None
|
|
384
|
+
|
|
385
|
+
def calc_loss_group(
|
|
386
|
+
self, group: NDArray[numpy.integer], sim: ReconsState, state: None
|
|
387
|
+
) -> t.Tuple[Float, None]:
|
|
388
|
+
xp = get_array_module(sim.probe.data)
|
|
389
|
+
probe_recip = xp.fft.fftshift(fft2(sim.probe.data), axes=(-1, -2))
|
|
390
|
+
|
|
391
|
+
g_y, g_x = img_grad(probe_recip)
|
|
392
|
+
cost = xp.sum(xp.sqrt(abs2(g_y) + abs2(g_x) + self.eps))
|
|
393
|
+
cost_scale = 1.0
|
|
394
|
+
|
|
395
|
+
return (cost * cost_scale * self.cost, state)
|
|
396
|
+
|
|
397
|
+
|
|
398
|
+
def img_grad(img: numpy.ndarray) -> t.Tuple[numpy.ndarray, numpy.ndarray]:
|
|
399
|
+
xp = get_array_module(img)
|
|
400
|
+
return (
|
|
401
|
+
xp.diff(img, axis=-2, append=img[..., -1:, :]),
|
|
402
|
+
xp.diff(img, axis=-1, append=img[..., :, -1:]),
|
|
403
|
+
)
|
|
@@ -0,0 +1,270 @@
|
|
|
1
|
+
import collections
|
|
2
|
+
import logging
|
|
3
|
+
import typing as t
|
|
4
|
+
|
|
5
|
+
import numpy
|
|
6
|
+
from numpy.typing import NDArray, DTypeLike
|
|
7
|
+
from typing_extensions import Self
|
|
8
|
+
|
|
9
|
+
from phaser.utils.num import (
|
|
10
|
+
get_array_module, to_real_dtype, to_complex_dtype,
|
|
11
|
+
fft2, ifft2, is_jax, to_numpy, block_until_ready,
|
|
12
|
+
)
|
|
13
|
+
from phaser.utils.misc import FloatKey, jax_dataclass, create_compact_groupings, create_sparse_groupings, shuffled
|
|
14
|
+
from phaser.utils.optics import fresnel_propagator, fourier_shift_filter
|
|
15
|
+
from phaser.state import ReconsState
|
|
16
|
+
from phaser.hooks.solver import NoiseModel
|
|
17
|
+
from phaser.hooks.regularization import GroupConstraint, IterConstraint, StateT
|
|
18
|
+
|
|
19
|
+
logger = logging.getLogger(__name__)
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class GroupManager:
|
|
23
|
+
def __init__(
|
|
24
|
+
self,
|
|
25
|
+
scan: NDArray[numpy.floating],
|
|
26
|
+
grouping: t.Optional[int] = None,
|
|
27
|
+
compact: bool = False,
|
|
28
|
+
seed: t.Any = None,
|
|
29
|
+
):
|
|
30
|
+
self.grouping = grouping or 64
|
|
31
|
+
self.compact = compact
|
|
32
|
+
self.seed = seed
|
|
33
|
+
self.groups: t.Optional[t.List[NDArray[numpy.int64]]] = None
|
|
34
|
+
self.n_groups: int = int(numpy.ceil(numpy.prod(scan.shape[:-1]) / self.grouping))
|
|
35
|
+
|
|
36
|
+
def _make(self, scan: NDArray[numpy.floating], i: int = 0) -> t.List[NDArray[numpy.int64]]:
|
|
37
|
+
if self.compact:
|
|
38
|
+
return create_compact_groupings(scan, self.grouping, seed=self.seed, i=i)
|
|
39
|
+
else:
|
|
40
|
+
return create_sparse_groupings(scan, self.grouping, seed=self.seed, i=i)
|
|
41
|
+
|
|
42
|
+
def iter(
|
|
43
|
+
self, scan: NDArray[numpy.floating],
|
|
44
|
+
i: int = 0, shuffle_groups: bool = False,
|
|
45
|
+
) -> t.Iterator[NDArray[numpy.int64]]:
|
|
46
|
+
if shuffle_groups or self.groups is None:
|
|
47
|
+
self.groups = self._make(scan, i)
|
|
48
|
+
return iter(self.groups)
|
|
49
|
+
# shuffle order of groups (though not the groups themselves)
|
|
50
|
+
return shuffled(self.groups, seed=self.seed, i=i)
|
|
51
|
+
|
|
52
|
+
def __len__(self) -> int:
|
|
53
|
+
return self.n_groups
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def stream_patterns(
|
|
57
|
+
groups: t.Iterable[NDArray[numpy.int64]], patterns: NDArray[numpy.floating],
|
|
58
|
+
xp: t.Any, buf_n: int = 1
|
|
59
|
+
) -> t.Iterator[t.Tuple[NDArray[numpy.int64], NDArray[numpy.floating]]]:
|
|
60
|
+
if buf_n == 0:
|
|
61
|
+
for group in groups:
|
|
62
|
+
group_patterns = xp.asarray(patterns[tuple(group)])
|
|
63
|
+
yield group, block_until_ready(group_patterns)
|
|
64
|
+
return
|
|
65
|
+
|
|
66
|
+
buf = collections.deque()
|
|
67
|
+
it = iter(groups)
|
|
68
|
+
|
|
69
|
+
for group in it:
|
|
70
|
+
buf.append((group, xp.asarray(patterns[tuple(group)])))
|
|
71
|
+
if len(buf) >= buf_n:
|
|
72
|
+
break
|
|
73
|
+
|
|
74
|
+
while len(buf) > 0:
|
|
75
|
+
(group, group_patterns) = buf.popleft()
|
|
76
|
+
yield group, block_until_ready(group_patterns)
|
|
77
|
+
|
|
78
|
+
# attempt to feed queue
|
|
79
|
+
try:
|
|
80
|
+
group = next(it)
|
|
81
|
+
buf.append((group, xp.asarray(patterns[tuple(group)])))
|
|
82
|
+
except StopIteration:
|
|
83
|
+
continue
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
@jax_dataclass(init=False, static_fields=('xp', 'dtype', 'noise_model', 'group_constraints', 'iter_constraints'), drop_fields=('ky', 'kx'))
|
|
87
|
+
class SimulationState:
|
|
88
|
+
state: ReconsState
|
|
89
|
+
|
|
90
|
+
ky: NDArray[numpy.floating]
|
|
91
|
+
kx: NDArray[numpy.floating]
|
|
92
|
+
|
|
93
|
+
noise_model: NoiseModel
|
|
94
|
+
group_constraints: t.Tuple[GroupConstraint[t.Any], ...]
|
|
95
|
+
iter_constraints: t.Tuple[IterConstraint[t.Any], ...]
|
|
96
|
+
|
|
97
|
+
noise_model_state: t.Any
|
|
98
|
+
group_constraint_states: t.Tuple[t.Any, ...]
|
|
99
|
+
iter_constraint_states: t.Tuple[t.Any, ...]
|
|
100
|
+
|
|
101
|
+
xp: t.Any
|
|
102
|
+
dtype: DTypeLike
|
|
103
|
+
start_iter: int
|
|
104
|
+
|
|
105
|
+
def __init__(
|
|
106
|
+
self, *,
|
|
107
|
+
state: ReconsState,
|
|
108
|
+
noise_model: NoiseModel[t.Any],
|
|
109
|
+
group_constraints: t.Tuple[GroupConstraint[t.Any], ...],
|
|
110
|
+
iter_constraints: t.Tuple[IterConstraint[t.Any], ...],
|
|
111
|
+
xp: t.Any,
|
|
112
|
+
dtype: DTypeLike,
|
|
113
|
+
noise_model_state: t.Optional[t.Any] = None,
|
|
114
|
+
group_constraint_states: t.Optional[t.Tuple[t.Any, ...]] = None,
|
|
115
|
+
iter_constraint_states: t.Optional[t.Tuple[t.Any, ...]] = None,
|
|
116
|
+
start_iter: t.Optional[int] = None,
|
|
117
|
+
):
|
|
118
|
+
self.xp = xp
|
|
119
|
+
self.dtype = dtype
|
|
120
|
+
self.state = state
|
|
121
|
+
|
|
122
|
+
self.noise_model = noise_model
|
|
123
|
+
self.group_constraints = group_constraints
|
|
124
|
+
self.iter_constraints = iter_constraints
|
|
125
|
+
|
|
126
|
+
self.noise_model_state = noise_model_state or noise_model.init_state(self.state)
|
|
127
|
+
self.group_constraint_states = group_constraint_states if group_constraint_states is not None else tuple(
|
|
128
|
+
reg.init_state(self.state) for reg in group_constraints
|
|
129
|
+
)
|
|
130
|
+
self.iter_constraint_states = iter_constraint_states if iter_constraint_states is not None else tuple(
|
|
131
|
+
reg.init_state(self.state) for reg in iter_constraints
|
|
132
|
+
)
|
|
133
|
+
|
|
134
|
+
self.start_iter = start_iter if start_iter is not None else self.state.iter.total_iter
|
|
135
|
+
(self.ky, self.kx) = state.probe.sampling.recip_grid(dtype=dtype, xp=xp)
|
|
136
|
+
|
|
137
|
+
def apply_group_constraints(self, group: NDArray[numpy.integer]) -> Self:
|
|
138
|
+
def apply_reg(reg: GroupConstraint[t.Any], state: t.Any):
|
|
139
|
+
(self.state, state) = reg.apply_group(group, self.state, state)
|
|
140
|
+
return state
|
|
141
|
+
|
|
142
|
+
self.group_constraint_states = tuple(
|
|
143
|
+
apply_reg(reg, state) for (reg, state) in zip(self.group_constraints, self.group_constraint_states)
|
|
144
|
+
)
|
|
145
|
+
|
|
146
|
+
return self
|
|
147
|
+
|
|
148
|
+
def apply_iter_constraints(self) -> Self:
|
|
149
|
+
def apply_reg(reg: IterConstraint[t.Any], state: t.Any):
|
|
150
|
+
(self.state, state) = reg.apply_iter(self.state, state)
|
|
151
|
+
return state
|
|
152
|
+
|
|
153
|
+
self.iter_constraint_states = tuple(
|
|
154
|
+
apply_reg(reg, state) for (reg, state) in zip(self.iter_constraints, self.iter_constraint_states)
|
|
155
|
+
)
|
|
156
|
+
|
|
157
|
+
return self
|
|
158
|
+
|
|
159
|
+
|
|
160
|
+
def make_propagators(state: ReconsState, bwlim_frac: t.Optional[float] = 2/3) -> t.Optional[NDArray[numpy.complexfloating]]:
|
|
161
|
+
xp = get_array_module(state.probe.data)
|
|
162
|
+
dtype = to_real_dtype(state.probe.data.dtype)
|
|
163
|
+
complex_dtype = to_complex_dtype(dtype)
|
|
164
|
+
|
|
165
|
+
(ky, kx) = state.probe.sampling.recip_grid(xp=xp, dtype=dtype)
|
|
166
|
+
|
|
167
|
+
# ignore last slice; we don't need it
|
|
168
|
+
delta_zs = to_numpy(state.object.thicknesses)[:-1]
|
|
169
|
+
if len(delta_zs) == 0:
|
|
170
|
+
return None
|
|
171
|
+
|
|
172
|
+
unique_zs = set(map(FloatKey, delta_zs))
|
|
173
|
+
|
|
174
|
+
if bwlim_frac is not None:
|
|
175
|
+
bwlim = numpy.min(state.probe.sampling.k_max) * bwlim_frac
|
|
176
|
+
k2 = ky**2 + kx**2
|
|
177
|
+
bwlim_mask = k2 <= bwlim**2
|
|
178
|
+
logger.info(f"Bandwidth limit: {bwlim * state.wavelength * 1e3:6.2f} mrad")
|
|
179
|
+
else:
|
|
180
|
+
bwlim_mask = xp.ones(ky.shape, dtype=numpy.bool_)
|
|
181
|
+
|
|
182
|
+
props = {
|
|
183
|
+
z: fresnel_propagator(ky, kx, state.wavelength, z).astype(complex_dtype) * bwlim_mask
|
|
184
|
+
for z in unique_zs
|
|
185
|
+
}
|
|
186
|
+
|
|
187
|
+
return xp.stack(
|
|
188
|
+
[props[FloatKey(z)] for z in delta_zs],
|
|
189
|
+
axis = 0
|
|
190
|
+
)
|
|
191
|
+
|
|
192
|
+
|
|
193
|
+
@t.overload
|
|
194
|
+
def cutout_group(
|
|
195
|
+
ky: NDArray[numpy.floating], kx: NDArray[numpy.floating],
|
|
196
|
+
state: ReconsState, group: NDArray[numpy.integer],
|
|
197
|
+
return_filters: t.Literal[False] = False
|
|
198
|
+
) -> t.Tuple[NDArray[numpy.complexfloating], NDArray[numpy.complexfloating], NDArray[numpy.floating]]:
|
|
199
|
+
...
|
|
200
|
+
|
|
201
|
+
@t.overload
|
|
202
|
+
def cutout_group(
|
|
203
|
+
ky: NDArray[numpy.floating], kx: NDArray[numpy.floating],
|
|
204
|
+
state: ReconsState, group: NDArray[numpy.integer],
|
|
205
|
+
return_filters: t.Literal[True]
|
|
206
|
+
) -> t.Tuple[NDArray[numpy.complexfloating], NDArray[numpy.complexfloating], NDArray[numpy.floating], NDArray[numpy.complexfloating]]:
|
|
207
|
+
...
|
|
208
|
+
|
|
209
|
+
def cutout_group(
|
|
210
|
+
ky: NDArray[numpy.floating], kx: NDArray[numpy.floating],
|
|
211
|
+
state: ReconsState, group: NDArray[numpy.integer],
|
|
212
|
+
return_filters: bool = False
|
|
213
|
+
):
|
|
214
|
+
"""Returns (probe, obj) in the cutout region"""
|
|
215
|
+
probes = state.probe.data
|
|
216
|
+
|
|
217
|
+
group_scan = state.scan[tuple(group)]
|
|
218
|
+
group_obj = state.object.sampling.get_view_at_pos(state.object.data, group_scan, probes.shape[-2:])
|
|
219
|
+
# group probes in real space
|
|
220
|
+
# shape (len(group), 1, Ny, Nx)
|
|
221
|
+
group_subpx_filters = fourier_shift_filter(ky, kx, state.object.sampling.get_subpx_shifts(group_scan, probes.shape[-2:]))[:, None, ...]
|
|
222
|
+
# shape (len(group), probe modes, Ny, Nx)
|
|
223
|
+
shifted_probes = ifft2(fft2(probes) * group_subpx_filters)
|
|
224
|
+
|
|
225
|
+
if return_filters:
|
|
226
|
+
return (shifted_probes, group_obj, group_scan, group_subpx_filters)
|
|
227
|
+
|
|
228
|
+
return (shifted_probes, group_obj, group_scan)
|
|
229
|
+
|
|
230
|
+
|
|
231
|
+
def slice_forwards(
|
|
232
|
+
props: t.Optional[NDArray[numpy.complexfloating]],
|
|
233
|
+
state: StateT,
|
|
234
|
+
f: t.Callable[[int, t.Optional[NDArray[numpy.complexfloating]], StateT], StateT]
|
|
235
|
+
) -> StateT:
|
|
236
|
+
if props is None:
|
|
237
|
+
return f(0, None, state)
|
|
238
|
+
|
|
239
|
+
n_slices = len(props) + 1
|
|
240
|
+
|
|
241
|
+
if is_jax(props):
|
|
242
|
+
import jax
|
|
243
|
+
state = jax.lax.fori_loop(0, n_slices - 1, lambda slice_i, state: f(slice_i, props[slice_i], state), state, unroll=False)
|
|
244
|
+
return f(n_slices - 1, None, state)
|
|
245
|
+
|
|
246
|
+
for slice_i in range(n_slices - 1):
|
|
247
|
+
state = f(slice_i, props[slice_i], state)
|
|
248
|
+
|
|
249
|
+
return f(n_slices - 1, None, state)
|
|
250
|
+
|
|
251
|
+
|
|
252
|
+
def slice_backwards(
|
|
253
|
+
props: t.Optional[NDArray[numpy.complexfloating]],
|
|
254
|
+
state: StateT,
|
|
255
|
+
f: t.Callable[[int, t.Optional[NDArray[numpy.complexfloating]], StateT], StateT]
|
|
256
|
+
) -> StateT:
|
|
257
|
+
if props is None:
|
|
258
|
+
return f(0, None, state)
|
|
259
|
+
|
|
260
|
+
n_slices = len(props) + 1
|
|
261
|
+
|
|
262
|
+
if is_jax(props):
|
|
263
|
+
import jax
|
|
264
|
+
state = jax.lax.fori_loop(1, n_slices, lambda i, state: f(n_slices - i, props[n_slices - i - 1], state), state, unroll=False)
|
|
265
|
+
return f(0, None, state)
|
|
266
|
+
|
|
267
|
+
for slice_i in range(n_slices - 1, 0, -1):
|
|
268
|
+
state = f(slice_i, props[slice_i - 1], state)
|
|
269
|
+
|
|
270
|
+
return f(0, None, state)
|
|
File without changes
|