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,142 @@
1
+ import logging
2
+ from pathlib import Path
3
+ import typing as t
4
+
5
+ import numpy
6
+
7
+ from phaser.utils.misc import mask_fraction_of_groups
8
+ from phaser.utils.num import cast_array_module, to_numpy, to_complex_dtype
9
+ from phaser.utils.io import OutputDir
10
+ from phaser.execute import Observer
11
+ from phaser.hooks import EngineArgs
12
+ from phaser.plan import ConventionalEnginePlan
13
+ from phaser.state import ReconsState
14
+ from phaser.types import process_flag, flag_any_true
15
+ from ..common.output import output_images, output_state
16
+ from ..common.simulation import SimulationState, make_propagators, GroupManager
17
+
18
+
19
+ def run_engine(args: EngineArgs, props: ConventionalEnginePlan) -> ReconsState:
20
+ logger = logging.getLogger(__name__)
21
+
22
+ xp = cast_array_module(args['xp'])
23
+ dtype = args['dtype']
24
+ observer: Observer = args.get('observer', [])
25
+ recons_name = args['recons_name']
26
+ engine_i = args['engine_i']
27
+ seed = args['seed']
28
+
29
+ logger.info(f"Starting engine #{args['engine_i'] + 1}...")
30
+
31
+ noise_model = props.noise_model(None)
32
+ group_constraints = tuple(reg(None) for reg in props.group_constraints)
33
+ iter_constraints = tuple(reg(None) for reg in props.iter_constraints)
34
+
35
+ update_probe = process_flag(props.update_probe)
36
+ update_object = process_flag(props.update_object)
37
+ update_positions = process_flag(props.update_positions)
38
+ calc_error = process_flag(props.calc_error)
39
+ save = process_flag(props.save)
40
+ save_images = process_flag(props.save_images)
41
+ # shuffle_groups defaults to True for sparse groups, False for compact groups
42
+ shuffle_groups = process_flag(props.shuffle_groups or not props.compact)
43
+
44
+ sim = SimulationState(
45
+ state=args['state'], noise_model=noise_model,
46
+ group_constraints=group_constraints, iter_constraints=iter_constraints,
47
+ xp=xp, dtype=dtype
48
+ )
49
+ patterns = args['data'].patterns
50
+ pattern_mask = xp.array(args['data'].pattern_mask)
51
+
52
+ assert patterns.dtype == sim.dtype
53
+ assert pattern_mask.dtype == sim.dtype
54
+ assert sim.state.object.data.dtype == to_complex_dtype(sim.dtype)
55
+ assert sim.state.probe.data.dtype == to_complex_dtype(sim.dtype)
56
+
57
+ solver = props.solver(props)
58
+ sim = solver.init(sim)
59
+ groups = GroupManager(sim.state.scan, props.grouping, props.compact, seed=seed)
60
+
61
+ any_output = flag_any_true(save, props.niter) or flag_any_true(save_images, props.niter)
62
+
63
+ with OutputDir(
64
+ props.save_options.out_dir, any_output,
65
+ engine_i=engine_i, name=recons_name,
66
+ group=groups.grouping, niter=props.niter,
67
+ solver=solver.name(),
68
+ noise_model=noise_model.name(),
69
+ ) as out_dir:
70
+ calc_error_mask = mask_fraction_of_groups(len(groups), props.calc_error_fraction)
71
+
72
+ position_solver = None if props.position_solver is None else props.position_solver(None)
73
+ position_solver_state = None if position_solver is None else position_solver.init_state(sim.state)
74
+
75
+ propagators = make_propagators(sim.state, props.bwlim_frac)
76
+
77
+ start_i = sim.state.iter.total_iter
78
+ observer.init_solver(sim.state, engine_i)
79
+
80
+ # runs rescaling
81
+ sim = solver.presolve(
82
+ sim, groups.iter(sim.state.scan),
83
+ patterns=patterns, pattern_mask=pattern_mask,
84
+ propagators=propagators
85
+ )
86
+
87
+ observer.start_solver()
88
+
89
+ for i in range(1, props.niter+1):
90
+ iter_update_positions = update_positions({'state': sim.state, 'niter': props.niter})
91
+ iter_shuffle_groups = shuffle_groups({'state': sim.state, 'niter': props.niter})
92
+
93
+ sim, pos_update, group_errors = solver.run_iteration(
94
+ sim, groups.iter(sim.state.scan, i, iter_shuffle_groups),
95
+ patterns=patterns, pattern_mask=pattern_mask, propagators=propagators,
96
+ update_object=update_object({'state': sim.state, 'niter': props.niter}),
97
+ update_probe=update_probe({'state': sim.state, 'niter': props.niter}),
98
+ update_positions=iter_update_positions,
99
+ calc_error=calc_error({'state': sim.state, 'niter': props.niter}),
100
+ calc_error_mask=calc_error_mask,
101
+ observer=observer,
102
+ )
103
+ assert sim.state.object.data.dtype == to_complex_dtype(sim.dtype)
104
+ assert sim.state.probe.data.dtype == to_complex_dtype(sim.dtype)
105
+
106
+ sim = sim.apply_iter_constraints()
107
+
108
+ if iter_update_positions:
109
+ if not position_solver:
110
+ raise ValueError("Updating positions with no PositionSolver specified")
111
+
112
+ # subtract mean position update
113
+ pos_update -= xp.mean(pos_update, tuple(range(pos_update.ndim - 1)))
114
+ pos_update, position_solver_state = position_solver.perform_update(sim.state.scan, pos_update, position_solver_state)
115
+ # subtract mean again (this can change with momentum)
116
+ pos_update -= xp.mean(pos_update, tuple(range(pos_update.ndim - 1)))
117
+ update_mag = xp.linalg.norm(pos_update, axis=-1, keepdims=True)
118
+ logger.info(f"Position update: mean {xp.mean(update_mag)}")
119
+ sim.state.scan += pos_update
120
+ assert sim.state.scan.dtype == sim.dtype
121
+
122
+ # check positions are at least overlapping object
123
+ sim.state.object.sampling.check_scan(sim.state.scan, sim.state.probe.sampling.extent / 2.)
124
+
125
+ error = None
126
+ if group_errors is not None:
127
+ error = float(to_numpy(xp.nanmean(xp.concatenate(group_errors))))
128
+
129
+ # TODO don't do this
130
+ sim.state.progress.iters = numpy.concatenate([sim.state.progress.iters, [i + start_i]])
131
+ sim.state.progress.detector_errors = numpy.concatenate([sim.state.progress.detector_errors, [error]])
132
+
133
+ observer.update_iteration(sim.state, i, props.niter, error)
134
+
135
+ if save({'state': sim.state, 'niter': props.niter}):
136
+ output_state(sim.state, out_dir, props.save_options)
137
+
138
+ if save_images({'state': sim.state, 'niter': props.niter}):
139
+ output_images(sim.state, out_dir, props.save_options)
140
+
141
+ observer.finish_solver()
142
+ return sim.state
@@ -0,0 +1,476 @@
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 cast_array_module, at, abs2, fft2, ifft2, jit, check_finite, to_complex_dtype, to_numpy
9
+ from phaser.hooks.solver import ConventionalSolver
10
+ from phaser.types import process_schedule
11
+ from phaser.plan import ConventionalEnginePlan, LSQMLSolverPlan, EPIESolverPlan
12
+ from phaser.execute import Observer
13
+ from phaser.engines.common.simulation import (
14
+ stream_patterns, SimulationState, cutout_group, slice_forwards, slice_backwards
15
+ )
16
+
17
+
18
+ class LSQMLSolver(ConventionalSolver):
19
+ def __init__(self, plan: ConventionalEnginePlan, props: LSQMLSolverPlan):
20
+ self.plan: LSQMLSolverPlan = props
21
+ self.engine_plan: ConventionalEnginePlan = plan
22
+
23
+ @classmethod
24
+ def name(cls) -> str:
25
+ return "LSQML"
26
+
27
+ def init(self, sim: SimulationState) -> SimulationState:
28
+ self.logger = logging.getLogger(__name__)
29
+ xp = sim.xp
30
+
31
+ self.obj_mag: NDArray[numpy.floating] = xp.zeros(sim.state.probe.data.shape[-2:], dtype=sim.dtype)
32
+ self.probe_mag: NDArray[numpy.floating] = xp.zeros_like(sim.state.object.data, dtype=sim.dtype)
33
+
34
+ return sim
35
+
36
+ def presolve(
37
+ self,
38
+ sim: SimulationState,
39
+ groups: t.Iterator[NDArray[numpy.int_]], *,
40
+ patterns: NDArray[numpy.floating],
41
+ pattern_mask: NDArray[numpy.floating],
42
+ propagators: t.Optional[NDArray[numpy.complexfloating]],
43
+ ) -> SimulationState:
44
+ rescale_factors = []
45
+
46
+ # precompute obj_mag, probe_mag, and rescale probe intensity
47
+ for (group, group_patterns) in stream_patterns(groups, patterns, xp=sim.xp, buf_n=self.engine_plan.buffer_n_groups):
48
+ (self.obj_mag, self.probe_mag, group_rescale_factors) = lsqml_dry_run(
49
+ sim, group, group_patterns, props=propagators, pattern_mask=pattern_mask,
50
+ obj_mag=self.obj_mag, probe_mag=self.probe_mag
51
+ )
52
+
53
+ rescale_factors.append(to_numpy(group_rescale_factors))
54
+
55
+ rescale_factors = numpy.concatenate(rescale_factors, axis=0)
56
+ rescale_factor = numpy.mean(rescale_factors)
57
+
58
+ self.logger.info("Pre-calculated intensities")
59
+ self.logger.info(f"Rescaling initial probe intensity by {rescale_factor:.2e}")
60
+ sim.state.probe.data *= numpy.sqrt(rescale_factor)
61
+ self.probe_mag *= rescale_factor
62
+
63
+ return sim
64
+
65
+ def run_iteration(
66
+ self,
67
+ sim: SimulationState,
68
+ groups: t.Iterator[NDArray[numpy.int_]], *,
69
+ patterns: NDArray[numpy.floating],
70
+ pattern_mask: NDArray[numpy.floating],
71
+ propagators: t.Optional[NDArray[numpy.complexfloating]],
72
+ update_object: bool,
73
+ update_probe: bool,
74
+ update_positions: bool,
75
+ calc_error: bool,
76
+ calc_error_mask: NDArray[numpy.bool_],
77
+ observer: 'Observer',
78
+ ) -> t.Tuple[SimulationState, NDArray[numpy.floating], t.List[NDArray[numpy.floating]]]:
79
+ xp = sim.xp
80
+
81
+ beta_object = process_schedule(self.plan.beta_object)({'state': sim.state, 'niter': self.engine_plan.niter})
82
+ beta_probe = process_schedule(self.plan.beta_probe)({'state': sim.state, 'niter': self.engine_plan.niter})
83
+ illum_reg_object = process_schedule(self.plan.illum_reg_object)({'state': sim.state, 'niter': self.engine_plan.niter})
84
+ illum_reg_probe = process_schedule(self.plan.illum_reg_probe)({'state': sim.state, 'niter': self.engine_plan.niter})
85
+ gamma = process_schedule(self.plan.gamma)({'state': sim.state, 'niter': self.engine_plan.niter})
86
+
87
+ new_obj_mag = xp.zeros_like(self.obj_mag)
88
+ new_probe_mag = xp.zeros_like(self.probe_mag)
89
+ pos_update = xp.zeros_like(sim.state.scan, dtype=sim.dtype)
90
+ iter_errors = []
91
+
92
+ for (group_i, (group, group_patterns)) in enumerate(stream_patterns(groups, patterns, xp=xp,
93
+ buf_n=self.engine_plan.buffer_n_groups)):
94
+ group_calc_error = calc_error and calc_error_mask[group_i]
95
+
96
+ (sim, new_obj_mag, new_probe_mag, errors, group_pos_update) = lsqml_run(
97
+ sim, group, group_patterns, pattern_mask=pattern_mask, props=propagators,
98
+ obj_mag=self.obj_mag, probe_mag=self.probe_mag,
99
+ new_obj_mag=new_obj_mag, new_probe_mag=new_probe_mag,
100
+ beta_object=beta_object, beta_probe=beta_probe,
101
+ update_object=update_object,
102
+ update_probe=update_probe,
103
+ update_position=update_positions,
104
+ calc_error=group_calc_error,
105
+ illum_reg_object=illum_reg_object,
106
+ illum_reg_probe=illum_reg_probe,
107
+ gamma=gamma,
108
+ )
109
+ check_finite(sim.state.object.data, sim.state.probe.data, context=f"object or probe, group {group_i}")
110
+ assert sim.state.object.data.dtype == to_complex_dtype(sim.dtype)
111
+ assert sim.state.probe.data.dtype == to_complex_dtype(sim.dtype)
112
+
113
+ sim = sim.apply_group_constraints(group)
114
+
115
+ if update_positions:
116
+ assert group_pos_update is not None
117
+ pos_update = at(pos_update, tuple(group)).set(group_pos_update)
118
+
119
+ observer.update_group(sim.state, self.engine_plan.send_every_group)
120
+
121
+ if group_calc_error:
122
+ assert errors is not None
123
+ iter_errors.append(errors)
124
+
125
+ self.obj_mag = new_obj_mag
126
+ self.probe_mag = new_probe_mag
127
+
128
+ return (sim, pos_update, iter_errors)
129
+
130
+
131
+ @partial(jit, donate_argnames=('obj_mag', 'probe_mag'))
132
+ def lsqml_dry_run(
133
+ sim: SimulationState,
134
+ group: NDArray[numpy.integer],
135
+ group_patterns: NDArray[numpy.floating], *,
136
+ pattern_mask: NDArray[numpy.floating],
137
+ props: t.Optional[NDArray[numpy.complexfloating]],
138
+ obj_mag: NDArray[numpy.floating],
139
+ probe_mag: NDArray[numpy.floating]
140
+ ) -> t.Tuple[NDArray[numpy.floating], NDArray[numpy.floating], NDArray[numpy.floating]]:
141
+ xp = cast_array_module(sim.xp)
142
+ (psi, group_obj, group_scan) = cutout_group(sim.ky, sim.kx, sim.state, group)
143
+
144
+ obj_mag += xp.sum(abs2(xp.prod(group_obj, axis=1)), axis=0)
145
+ obj_grid = sim.state.object.sampling
146
+
147
+ def run_slice(slice_i: int, prop: t.Optional[NDArray[numpy.complexfloating]], state):
148
+ (probe_mag, psi) = state
149
+ probe_mag = at(probe_mag, slice_i).set(
150
+ obj_grid.add_view_at_pos(probe_mag[slice_i], group_scan, xp.sum(abs2(psi), axis=1))
151
+ )
152
+
153
+ if prop is not None:
154
+ psi = ifft2(fft2(psi * group_obj[:, slice_i, None]) * prop)
155
+
156
+ return (probe_mag, psi)
157
+
158
+ (probe_mag, psi) = slice_forwards(props, (probe_mag, psi), run_slice)
159
+
160
+ # modeled and experimental intensity
161
+ # summed over incoherent modes and over the pattern
162
+ model_intensity = xp.sum(abs2(fft2(psi)), axis=(1, -2, -1))
163
+ exp_intensity = xp.sum(group_patterns * pattern_mask, axis=(-2, -1))
164
+
165
+ return (obj_mag, probe_mag, exp_intensity / model_intensity)
166
+
167
+ # TODO: pass LSQMLSolverPlan in here for parameters
168
+
169
+ @partial(
170
+ jit,
171
+ donate_argnames=('sim', 'new_obj_mag', 'new_probe_mag'),
172
+ static_argnames=('update_object', 'update_probe', 'update_position', 'calc_error'),
173
+ )
174
+ def lsqml_run(
175
+ sim: SimulationState,
176
+ group: NDArray[numpy.integer],
177
+ group_patterns: NDArray[numpy.floating], *,
178
+ pattern_mask: NDArray[numpy.floating],
179
+ props: t.Optional[NDArray[numpy.complexfloating]],
180
+ obj_mag: NDArray[numpy.floating],
181
+ probe_mag: NDArray[numpy.floating],
182
+ new_obj_mag: NDArray[numpy.floating],
183
+ new_probe_mag: NDArray[numpy.floating],
184
+ beta_object: float = 0.9,
185
+ beta_probe: float = 0.9,
186
+ update_object: bool = True,
187
+ update_probe: bool = True,
188
+ update_position: bool = True,
189
+ calc_error: bool = True,
190
+ illum_reg_object: float,
191
+ illum_reg_probe: float,
192
+ gamma: float,
193
+ ) -> t.Tuple[SimulationState, NDArray[numpy.floating], NDArray[numpy.floating], t.Optional[NDArray[numpy.floating]], t.Optional[NDArray[numpy.floating]]]:
194
+ xp = cast_array_module(sim.xp)
195
+ obj_grid = sim.state.object.sampling
196
+ n_slices = sim.state.object.data.shape[0]
197
+
198
+ eps = 1e-16
199
+
200
+ (probes, group_obj, group_scan, subpx_filters) = cutout_group(sim.ky, sim.kx, sim.state, group, return_filters=True)
201
+ psi = xp.zeros((n_slices, *probes.shape), dtype=probes.dtype)
202
+ psi = at(psi, 0).set(probes)
203
+
204
+ group_probe_mag = xp.zeros_like(probe_mag)
205
+ #group_obj_mag = xp.sum(abs2(group_obj[:, 0]), axis=0)
206
+ group_obj_mag = xp.sum(abs2(xp.prod(group_obj, axis=1)), axis=0)
207
+
208
+ def sim_slice(slice_i: int, prop: t.Optional[NDArray[numpy.complexfloating]], state):
209
+ (group_probe_mag, psi) = state
210
+
211
+ group_probe_mag = at(group_probe_mag, slice_i).set(
212
+ obj_grid.add_view_at_pos(group_probe_mag[slice_i], group_scan, xp.sum(abs2(psi[slice_i]), axis=1))
213
+ )
214
+
215
+ if prop is not None:
216
+ psi = at(psi, slice_i + 1).set(
217
+ ifft2(fft2(psi[slice_i] * group_obj[:, slice_i, None]) * prop)
218
+ )
219
+
220
+ return (group_probe_mag, psi)
221
+
222
+ (group_probe_mag, psi) = slice_forwards(props, (group_probe_mag, psi), sim_slice)
223
+
224
+ new_obj_mag += group_obj_mag
225
+ new_probe_mag += group_probe_mag
226
+
227
+ model_wave = fft2(psi[-1] * group_obj[:, -1, None])
228
+ # sum over incoherent modes
229
+ model_intensity = xp.sum(abs2(model_wave), axis=1, keepdims=True)
230
+ # experimental data
231
+ # group_patterns = xp.array(sim.patterns[tuple(group)])[:, None]
232
+
233
+ errors = xp.sqrt(xp.nansum((model_intensity - group_patterns[:, None])**2, axis=(1, -1, -2))) if calc_error else None
234
+
235
+ (chi, sim.noise_model_state) = sim.noise_model.calc_wave_update(model_wave, model_intensity, group_patterns[:, None], pattern_mask, sim.noise_model_state)
236
+ chi = ifft2(chi)
237
+
238
+ def update_slice(slice_i: int, prop: t.Optional[NDArray[numpy.complexfloating]], state):
239
+ (sim, chi) = state
240
+
241
+ delta_P = chi * xp.conj(group_obj[:, slice_i, None])
242
+
243
+ if update_object:
244
+ delta_O = chi * xp.conj(psi[slice_i])
245
+ alpha_O = xp.sum(xp.sum(xp.real(chi * xp.conj(delta_O * psi[slice_i])), axis=(-1, -2), keepdims=True), axis=1) / (xp.sum(abs2(delta_O * psi[slice_i])) + gamma)
246
+
247
+ # average object update
248
+ delta_O_avg = xp.zeros_like(sim.state.object.data[0])
249
+ delta_O_avg = obj_grid.add_view_at_pos(delta_O_avg, group_scan, xp.sum(delta_O, axis=1))
250
+ delta_O_avg /= (group_probe_mag[slice_i] + illum_reg_object)
251
+
252
+ obj_update = beta_object * xp.sum(alpha_O * delta_O_avg * group_probe_mag[slice_i], axis=0) / (group_probe_mag[slice_i] + eps)
253
+ sim.state.object.data = at(sim.state.object.data, slice_i).add(obj_update)
254
+
255
+ if prop is not None:
256
+ chi = ifft2(fft2(delta_P) * prop.conj())
257
+ elif update_probe:
258
+ delta_P_avg = ifft2(xp.sum(fft2(delta_P) * subpx_filters.conj(), axis=0))
259
+ delta_P_avg /= (group_obj_mag + illum_reg_probe)
260
+
261
+ # update step per probe mode
262
+ alpha_P = xp.sum(xp.real(chi * xp.conj(delta_P * group_obj[:, slice_i, None])), axis=(-1, -2), keepdims=True) / (xp.sum(abs2(delta_P * group_obj[:, slice_i, None])) + gamma)
263
+
264
+ probe_update = beta_probe * xp.sum(alpha_P * delta_P_avg * group_obj_mag, axis=0) / (group_obj_mag + eps)
265
+ sim.state.probe.data += probe_update
266
+
267
+ return (sim, chi)
268
+
269
+ (sim, chi) = slice_backwards(props, (sim, chi), update_slice)
270
+
271
+ if update_position:
272
+ def calc_pos_step(probes_fft: NDArray[numpy.complexfloating], kx: NDArray[numpy.floating]) -> NDArray[numpy.floating]:
273
+ delta_P_x = ifft2(probes_fft * -2.j*numpy.pi * kx)
274
+
275
+ prod = delta_P_x * group_obj[:, 0, None]
276
+ alpha = xp.sum(xp.real(chi * xp.conj(prod)), axis=(1, -1, -2)) / xp.sum(abs2(prod))
277
+ return alpha
278
+
279
+ # update directions
280
+ probes_fft = fft2(probes)
281
+ probes_fft /= xp.sum(abs2(probes), axis=(1, -1, -2), keepdims=True)
282
+ pos_update = xp.stack(tuple(calc_pos_step(probes_fft, k) for k in (sim.ky, sim.kx)), axis=-1)
283
+ else:
284
+ pos_update = None
285
+
286
+ return (sim, new_obj_mag, new_probe_mag, errors, pos_update)
287
+
288
+
289
+ class EPIESolver(ConventionalSolver):
290
+ def __init__(self, engine_plan: ConventionalEnginePlan, props: EPIESolverPlan):
291
+ self.engine_plan: ConventionalEnginePlan = engine_plan
292
+ self.plan: EPIESolverPlan = props
293
+
294
+ @classmethod
295
+ def name(cls) -> str:
296
+ return "ePIE"
297
+
298
+ def init(self, sim: SimulationState) -> SimulationState:
299
+ self.logger = logging.getLogger(__name__)
300
+ return sim
301
+
302
+ def presolve(
303
+ self,
304
+ sim: SimulationState,
305
+ groups: t.Iterator[NDArray[numpy.int_]], *,
306
+ patterns: NDArray[numpy.floating],
307
+ pattern_mask: NDArray[numpy.floating],
308
+ propagators: t.Optional[NDArray[numpy.complexfloating]],
309
+ ) -> SimulationState:
310
+ rescale_factors = []
311
+ for (group, group_patterns) in stream_patterns(groups, patterns, xp=sim.xp,
312
+ buf_n=self.engine_plan.buffer_n_groups):
313
+ group_rescale_factors = epie_dry_run(
314
+ sim, group, group_patterns, pattern_mask=pattern_mask, props=propagators
315
+ )
316
+ rescale_factors.append(to_numpy(group_rescale_factors))
317
+
318
+ rescale_factors = numpy.concatenate(rescale_factors, axis=0)
319
+ rescale_factor = numpy.mean(rescale_factors)
320
+
321
+ self.logger.info("Pre-calculated intensities")
322
+ self.logger.info(f"Rescaling initial probe intensity by {rescale_factor:.2e}")
323
+ sim.state.probe.data *= numpy.sqrt(rescale_factor)
324
+
325
+ return sim
326
+
327
+ def run_iteration(
328
+ self,
329
+ sim: SimulationState,
330
+ groups: t.Iterator[NDArray[numpy.int_]], *,
331
+ patterns: NDArray[numpy.floating],
332
+ pattern_mask: NDArray[numpy.floating],
333
+ propagators: t.Optional[NDArray[numpy.complexfloating]],
334
+ update_object: bool,
335
+ update_probe: bool,
336
+ update_positions: bool,
337
+ calc_error: bool,
338
+ calc_error_mask: NDArray[numpy.bool_],
339
+ observer: 'Observer',
340
+ ) -> t.Tuple[SimulationState, NDArray[numpy.floating], t.List[NDArray[numpy.floating]]]:
341
+ xp = sim.xp
342
+
343
+ # TODO: ePIE position update
344
+ pos_update = xp.zeros_like(sim.state.scan)
345
+ iter_errors = []
346
+
347
+ beta_object = process_schedule(self.plan.beta_object)({'state': sim.state, 'niter': self.engine_plan.niter})
348
+ beta_probe = process_schedule(self.plan.beta_probe)({'state': sim.state, 'niter': self.engine_plan.niter})
349
+
350
+ for (group_i, (group, group_patterns)) in enumerate(stream_patterns(groups, patterns, xp=xp,
351
+ buf_n=self.engine_plan.buffer_n_groups)):
352
+ group_calc_error = calc_error and calc_error_mask[group_i]
353
+
354
+ (sim, errors) = epie_run(
355
+ sim, group, group_patterns,
356
+ pattern_mask=pattern_mask,
357
+ props=propagators,
358
+ beta_object=beta_object,
359
+ beta_probe=beta_probe,
360
+ update_object=update_object,
361
+ update_probe=update_probe,
362
+ )
363
+ check_finite(sim.state.object.data, sim.state.probe.data, context=f"object or probe, group {group_i}")
364
+ assert sim.state.object.data.dtype == to_complex_dtype(sim.dtype)
365
+ assert sim.state.probe.data.dtype == to_complex_dtype(sim.dtype)
366
+
367
+ sim = sim.apply_group_constraints(group)
368
+
369
+ observer.update_group(sim.state, self.engine_plan.send_every_group)
370
+
371
+ if group_calc_error:
372
+ assert errors is not None
373
+ iter_errors.append(errors)
374
+
375
+ return (sim, pos_update, iter_errors)
376
+
377
+
378
+ @partial(jit)
379
+ def epie_dry_run(
380
+ sim: SimulationState,
381
+ group: NDArray[numpy.integer],
382
+ group_patterns: NDArray[numpy.floating], *,
383
+ pattern_mask: NDArray[numpy.floating],
384
+ props: t.Optional[NDArray[numpy.complexfloating]],
385
+ ) -> NDArray[numpy.floating]:
386
+ xp = cast_array_module(sim.xp)
387
+ (psi, group_obj, group_scan) = cutout_group(sim.ky, sim.kx, sim.state, group)
388
+
389
+ def run_slice(slice_i: int, prop: t.Optional[NDArray[numpy.complexfloating]], psi):
390
+ if prop is not None:
391
+ psi = ifft2(fft2(psi * group_obj[:, slice_i, None]) * prop)
392
+
393
+ return psi
394
+
395
+ psi = slice_forwards(props, psi, run_slice)
396
+
397
+ # modeled and experimental intensity
398
+ # summed over incoherent modes and over the pattern
399
+ model_intensity = xp.sum(abs2(fft2(psi)), axis=(1, -2, -1))
400
+ exp_intensity = xp.sum(group_patterns * pattern_mask, axis=(-2, -1))
401
+
402
+ return exp_intensity / model_intensity
403
+
404
+
405
+ @partial(jit, donate_argnames=('sim',), static_argnames=('update_object', 'update_probe', 'calc_error'))
406
+ def epie_run(
407
+ sim: SimulationState,
408
+ group: NDArray[numpy.integer],
409
+ group_patterns: NDArray[numpy.floating], *,
410
+ pattern_mask: NDArray[numpy.floating],
411
+ props: t.Optional[NDArray[numpy.complexfloating]],
412
+ beta_object: float = 0.9,
413
+ beta_probe: float = 0.9,
414
+ update_object: bool = True,
415
+ update_probe: bool = True,
416
+ calc_error: bool = True,
417
+ ) -> t.Tuple[SimulationState, t.Optional[NDArray[numpy.floating]]]:
418
+ xp = cast_array_module(sim.xp)
419
+ obj_grid = sim.state.object.sampling
420
+ n_slices = sim.state.object.data.shape[0]
421
+
422
+ (probes, group_obj, group_scan, subpx_filters) = cutout_group(sim.ky, sim.kx, sim.state, group, return_filters=True)
423
+ psi = xp.zeros((n_slices, *probes.shape), dtype=probes.dtype)
424
+ psi = at(psi, 0).set(probes)
425
+
426
+ def sim_slice(slice_i: int, prop: t.Optional[NDArray[numpy.complexfloating]], psi):
427
+ if prop is not None:
428
+ psi = at(psi, slice_i + 1).set(
429
+ ifft2(fft2(psi[slice_i] * group_obj[:, slice_i, None]) * prop)
430
+ )
431
+
432
+ return psi
433
+
434
+ psi = slice_forwards(props, psi, sim_slice)
435
+
436
+ model_wave = fft2(psi[-1] * group_obj[:, -1, None])
437
+ # sum over incoherent modes
438
+ model_intensity = xp.sum(abs2(model_wave), axis=1, keepdims=True)
439
+
440
+ errors = xp.sqrt(xp.nansum((model_intensity - group_patterns[:, None])**2, axis=(1, -1, -2))) if calc_error else None
441
+ (chi, sim.noise_model_state) = sim.noise_model.calc_wave_update(
442
+ model_wave, model_intensity, group_patterns[:, None], pattern_mask, sim.noise_model_state
443
+ )
444
+ chi = ifft2(chi)
445
+
446
+ def update_slice(slice_i: int, prop: t.Optional[NDArray[numpy.complexfloating]], state):
447
+ (sim, chi) = state
448
+
449
+ probe_update = group_obj[:, slice_i, None].conj() * chi / xp.max(
450
+ abs2(group_obj[:, slice_i, None]),
451
+ axis=(-1, -2), keepdims=True
452
+ )
453
+
454
+ if update_object:
455
+ # average incoherent modes
456
+ group_obj_update = beta_object/n_slices * xp.sum(psi[slice_i].conj() * chi, axis=1) / xp.max(
457
+ xp.sum(abs2(psi[slice_i]), axis=1),
458
+ axis=(-1, -2), keepdims=True
459
+ )
460
+
461
+ sim.state.object.data = at(sim.state.object.data, slice_i).set(
462
+ obj_grid.add_view_at_pos(sim.state.object.data[slice_i], group_scan, group_obj_update)
463
+ )
464
+
465
+ if prop is not None:
466
+ chi = ifft2(fft2(probe_update) * prop.conj())
467
+ elif update_probe:
468
+ # average probe updates in group
469
+ probe_update = ifft2(xp.mean(fft2(probe_update) * subpx_filters.conj(), axis=0))
470
+ sim.state.probe.data += beta_probe * probe_update
471
+
472
+ return (sim, chi)
473
+
474
+ (sim, chi) = slice_backwards(props, (sim, chi), update_slice)
475
+
476
+ return (sim, errors)