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.
Files changed (67) hide show
  1. phaser/__init__.py +0 -0
  2. phaser/__main__.py +5 -0
  3. phaser/engines/common/__init__.py +0 -0
  4. phaser/engines/common/noise_models.py +113 -0
  5. phaser/engines/common/output.py +189 -0
  6. phaser/engines/common/position_correction.py +62 -0
  7. phaser/engines/common/regularizers.py +403 -0
  8. phaser/engines/common/simulation.py +270 -0
  9. phaser/engines/conventional/__init__.py +0 -0
  10. phaser/engines/conventional/run.py +142 -0
  11. phaser/engines/conventional/solvers.py +476 -0
  12. phaser/engines/gradient/run.py +451 -0
  13. phaser/engines/gradient/solvers.py +139 -0
  14. phaser/execute.py +371 -0
  15. phaser/hooks/__init__.py +158 -0
  16. phaser/hooks/hook.py +159 -0
  17. phaser/hooks/io/empad.py +88 -0
  18. phaser/hooks/object.py +25 -0
  19. phaser/hooks/preprocessing.py +133 -0
  20. phaser/hooks/probe.py +24 -0
  21. phaser/hooks/regularization.py +97 -0
  22. phaser/hooks/scan.py +27 -0
  23. phaser/hooks/schedule.py +75 -0
  24. phaser/hooks/solver.py +169 -0
  25. phaser/io/__init__.py +0 -0
  26. phaser/io/empad.py +212 -0
  27. phaser/main.py +92 -0
  28. phaser/plan.py +184 -0
  29. phaser/py.typed +0 -0
  30. phaser/state.py +249 -0
  31. phaser/types.py +305 -0
  32. phaser/utils/__init__.py +0 -0
  33. phaser/utils/_cuda_kernels.py +213 -0
  34. phaser/utils/_jax_kernels.py +98 -0
  35. phaser/utils/analysis.py +263 -0
  36. phaser/utils/image.py +201 -0
  37. phaser/utils/io.py +402 -0
  38. phaser/utils/misc.py +295 -0
  39. phaser/utils/num.py +800 -0
  40. phaser/utils/object.py +578 -0
  41. phaser/utils/optics.py +377 -0
  42. phaser/utils/physics.py +88 -0
  43. phaser/utils/plotting.py +699 -0
  44. phaser/utils/scan.py +60 -0
  45. phaser/web/__init__.py +0 -0
  46. phaser/web/dist/03510a839ccb97b0da9f.module.wasm +0 -0
  47. phaser/web/dist/9573273f862f4f5d9644.module.wasm +0 -0
  48. phaser/web/dist/bundle-dashboard.js +712 -0
  49. phaser/web/dist/bundle-manager.js +210 -0
  50. phaser/web/dist/bundle-vendors-node_modules_wasm-array_wasm_array_js.js +106 -0
  51. phaser/web/dist/style.css +152 -0
  52. phaser/web/notebook.py +269 -0
  53. phaser/web/routes.py +211 -0
  54. phaser/web/server.py +540 -0
  55. phaser/web/slurm.py +180 -0
  56. phaser/web/templates/base.html +13 -0
  57. phaser/web/templates/dashboard.html +14 -0
  58. phaser/web/templates/manager.html +19 -0
  59. phaser/web/types.py +255 -0
  60. phaser/web/util.py +116 -0
  61. phaser/web/worker.py +200 -0
  62. phaserem-0.1.dist-info/METADATA +121 -0
  63. phaserem-0.1.dist-info/RECORD +67 -0
  64. phaserem-0.1.dist-info/WHEEL +5 -0
  65. phaserem-0.1.dist-info/entry_points.txt +2 -0
  66. phaserem-0.1.dist-info/licenses/LICENSE.txt +373 -0
  67. 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