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
phaser/execute.py ADDED
@@ -0,0 +1,371 @@
1
+ import dataclasses
2
+ import logging
3
+ import time
4
+ import typing as t
5
+
6
+ import numpy
7
+ import pane
8
+
9
+ from phaser.utils.num import cast_array_module, get_backend_module, xp_is_jax, Sampling, to_complex_dtype
10
+ from phaser.utils.object import ObjectSampling
11
+ from .hooks import Hook, ObjectHook, RawData
12
+ from .plan import GradientEnginePlan, ReconsPlan, EnginePlan, ScanHook, ProbeHook
13
+ from .state import Patterns, ReconsState, PartialReconsState, IterState, ProgressState
14
+
15
+
16
+ class Observer:
17
+ def __init__(self):
18
+ self.solver_start_time: t.Optional[float] = None
19
+ self.iter_start_time: t.Optional[float] = None
20
+ self.engine_i: int = 0
21
+ self.start_iter: int = 0
22
+
23
+ def _format_hhmmss(self, seconds: float) -> str:
24
+ hh, ss = divmod(seconds, (60 * 60))
25
+ mm, ss = divmod(ss, 60)
26
+ return f"{int(hh):02d}:{int(mm):02d}:{ss:06.3f}"
27
+
28
+ def _format_mmss(self, seconds: float) -> str:
29
+ mm, ss = divmod(seconds, 60)
30
+ return f"{int(mm):02d}:{ss:06.3f}"
31
+
32
+ def init_solver(self, init_state: ReconsState, engine_i: int):
33
+ self.engine_i = engine_i
34
+ self.start_iter = init_state.iter.total_iter
35
+
36
+ init_state.iter = IterState(self.engine_i, 1, self.start_iter + 1)
37
+
38
+ def start_solver(self):
39
+ logging.info("Engine initialized")
40
+ self.iter_start_time = self.solver_start_time = time.monotonic()
41
+
42
+ def heartbeat(self):
43
+ pass
44
+
45
+ def update_group(self, state: t.Union[ReconsState, PartialReconsState], force: bool = False):
46
+ pass
47
+
48
+ def update_iteration(self, state: t.Union[ReconsState, PartialReconsState], i: int, n: int, error: t.Optional[float] = None):
49
+ finish_time = time.monotonic()
50
+
51
+ if self.iter_start_time is not None:
52
+ delta = finish_time - self.iter_start_time
53
+ time_s = f" [{self._format_mmss(delta)}]"
54
+ else:
55
+ time_s = ""
56
+
57
+ w = len(str(n))
58
+
59
+ error_s = f" Error: {error:.3e}" if error is not None else ""
60
+ logging.info(f"Finished iter {i:{w}}/{n}{time_s}{error_s}")
61
+
62
+ state.iter = IterState(self.engine_i, i + 1, self.start_iter + i + 1)
63
+ self.iter_start_time = finish_time
64
+
65
+ def finish_solver(self):
66
+ logging.info("Solver finished!")
67
+ if self.solver_start_time is not None:
68
+ finish_time = time.monotonic()
69
+ delta = finish_time - self.solver_start_time
70
+ logging.info(f"Total time: {self._format_hhmmss(delta)}")
71
+
72
+
73
+ def execute_plan(plan: ReconsPlan, observer: t.Optional[Observer] = None):
74
+ xp = get_backend_module(plan.backend)
75
+
76
+ if observer is None:
77
+ observer = Observer()
78
+
79
+ patterns, state = initialize_reconstruction(plan, xp)
80
+ dtype = patterns.patterns.dtype
81
+
82
+ for (engine_i, engine) in enumerate(plan.engines):
83
+ logging.info(f"Preparing for engine #{engine_i + 1}...")
84
+ patterns, state = prepare_for_engine(patterns, state, xp, t.cast(EnginePlan, engine.props))
85
+ state = engine({
86
+ 'data': patterns,
87
+ 'state': state,
88
+ 'dtype': dtype,
89
+ 'xp': xp,
90
+ 'recons_name': plan.name,
91
+ 'engine_i': engine_i,
92
+ 'observer': observer,
93
+ 'seed': None,
94
+ })
95
+
96
+ logging.info("Reconstruction finished!")
97
+
98
+
99
+ def load_raw_data(
100
+ plan: ReconsPlan, xp: t.Any, seed: t.Any = None,
101
+ init_state: t.Union[ReconsState, PartialReconsState, None] = None
102
+ ) -> RawData:
103
+ dtype: type = numpy.float32 if plan.dtype == 'float32' else numpy.float64
104
+
105
+ raw_data = plan.raw_data(None)
106
+
107
+ wavelength = plan.wavelength or raw_data['wavelength']
108
+ if wavelength is None:
109
+ raise ValueError("`wavelength` must be specified by raw_data or manually")
110
+
111
+ if init_state is None:
112
+ init_state = PartialReconsState()
113
+
114
+ if init_state.wavelength is not None and not numpy.isclose(init_state.wavelength, wavelength):
115
+ logging.warning(f"Wavelength of reconstruction ({wavelength:.2e}) doesn't match wavelength " \
116
+ f"of previous state ({init_state.wavelength:.2e})")
117
+
118
+ raw_data['scan_hook'] = pane.into_data(merge( # type: ignore
119
+ pane.from_data(t.cast(dict, raw_data['scan_hook']), ScanHook) if raw_data['scan_hook'] is not None else None,
120
+ _MISSING if plan.init.scan in (None, {}) else plan.init.scan
121
+ ))
122
+ raw_data['probe_hook'] = pane.into_data(merge( # type: ignore
123
+ pane.from_data(t.cast(dict, raw_data['probe_hook']), ProbeHook) if raw_data['probe_hook'] is not None else None,
124
+ _MISSING if plan.init.probe in (None, {}) else plan.init.probe
125
+ ))
126
+ #print(f"scan_hook: {raw_data['scan_hook']}")
127
+ #print(f"probe_hook: {raw_data['probe_hook']}")
128
+
129
+ if raw_data['scan_hook'] is None and init_state.scan is None:
130
+ raise ValueError("`scan` must be specified by raw data, previous state, or manually in `init.scan`")
131
+ if raw_data['probe_hook'] is None and init_state.probe is None:
132
+ raise ValueError("`probe` must be specified by raw data, previous state, or manually in `init.probe`")
133
+ if raw_data['scan_hook'] == {}:
134
+ raise ValueError("Manual `init.scan` specified to override initial state, but scan was not provided by the raw data")
135
+ if raw_data['probe_hook'] == {}:
136
+ raise ValueError("Manual `init.probe` specified to override initial state, but probe was not provided by the raw data")
137
+
138
+ raw_data['wavelength'] = wavelength
139
+ raw_data['seed'] = seed
140
+
141
+ # normalize pattern intensity
142
+ #raw_data['patterns'] /= numpy.mean(numpy.sum(raw_data['patterns'], axis=(-1, -2)))
143
+ # ensure raw data is of the correct type
144
+ if raw_data['patterns'].dtype != dtype:
145
+ raw_data['patterns'] = raw_data['patterns'].astype(dtype)
146
+
147
+ # process post_load hooks:
148
+ for p in plan.post_load:
149
+ raw_data = p(raw_data)
150
+
151
+ # materialize memmap
152
+ if isinstance(raw_data['patterns'], numpy.memmap):
153
+ raw_data['patterns'] = raw_data['patterns'].copy()
154
+
155
+ return raw_data
156
+
157
+
158
+ def initialize_reconstruction(
159
+ plan: ReconsPlan, xp: t.Any, seed: t.Any = None,
160
+ init_state: t.Union[ReconsState, PartialReconsState, None] = None
161
+ ) -> t.Tuple[Patterns, ReconsState]:
162
+ xp = cast_array_module(xp)
163
+
164
+ logging.basicConfig(level=logging.INFO)
165
+
166
+ logging.info("Executing plan...")
167
+
168
+ dtype: t.Type[numpy.floating] = numpy.float32 if plan.dtype == 'float32' else numpy.float64
169
+ cdtype: t.Type[numpy.complexfloating] = to_complex_dtype(dtype)
170
+
171
+ logging.info(f"dtype: {dtype} array backend: {xp.__name__}")
172
+ if xp_is_jax(xp):
173
+ import jax
174
+ logging.info(f"jax backend: {jax.default_backend()} devices: {jax.devices()}")
175
+
176
+ if init_state is None:
177
+ if plan.init.state is not None:
178
+ path = plan.init.state.expanduser()
179
+ logging.info(f"Loading inital state from '{path}'...")
180
+ init_state = PartialReconsState.read_hdf5(path)
181
+ else:
182
+ init_state = PartialReconsState()
183
+
184
+ raw_data = load_raw_data(plan, xp, seed, init_state=init_state)
185
+
186
+ data = Patterns(raw_data['patterns'], raw_data['mask'])
187
+ sampling = raw_data['sampling']
188
+ wavelength = t.cast(float, raw_data['wavelength'])
189
+ probe_hook = raw_data['probe_hook']
190
+ scan_hook = raw_data['scan_hook']
191
+ del raw_data
192
+
193
+ if init_state.probe is not None and plan.init.probe is None:
194
+ logging.info("Re-using probe from initial state...")
195
+ probe = init_state.probe
196
+ probe.data = probe.data.astype(cdtype)
197
+
198
+ if probe.sampling != sampling:
199
+ logging.info("Resampling patterns to probe from initial state...")
200
+ data.patterns = sampling.resample_recip(data.patterns, probe.sampling)
201
+ data.pattern_mask = sampling.resample_recip(data.pattern_mask, probe.sampling)
202
+ sampling = probe.sampling
203
+
204
+ else:
205
+ logging.info("Initializing probe...")
206
+ probe = pane.from_data(probe_hook, ProbeHook)( # type: ignore
207
+ {'sampling': sampling, 'wavelength': wavelength, 'dtype': dtype, 'seed': seed, 'xp': xp}
208
+ )
209
+ if probe.data.ndim == 2:
210
+ probe.data = probe.data.reshape((1, *probe.data.shape))
211
+
212
+ if init_state.scan is not None and plan.init.scan is None:
213
+ logging.info("Re-using scan from initial state...")
214
+ scan = init_state.scan
215
+ else:
216
+ logging.info("Initializing scan...")
217
+ scan = pane.from_data(scan_hook, ScanHook)( # type: ignore
218
+ {'dtype': dtype, 'seed': seed, 'xp': xp}
219
+ )
220
+
221
+ obj_pad_px: float = plan.engines[0].obj_pad_px if len(plan.engines) > 0 else 5.0 # type: ignore
222
+ obj_sampling = ObjectSampling.from_scan(
223
+ scan, sampling.sampling, sampling.extent / 2. + obj_pad_px * sampling.sampling
224
+ )
225
+
226
+ if init_state.object is not None and plan.init.object is None:
227
+ logging.info("Re-using object from initial state...")
228
+ obj = init_state.object
229
+ obj.data = obj.data.astype(cdtype)
230
+ else:
231
+ logging.info("Initializing object...")
232
+ obj = (plan.init.object or pane.from_data('random', ObjectHook))({
233
+ 'sampling': obj_sampling, 'slices': plan.slices, 'wavelength': wavelength,
234
+ 'dtype': dtype, 'seed': seed, 'xp': xp
235
+ })
236
+ if obj.data.ndim == 2:
237
+ obj.data = obj.data.reshape((1, *obj.data.shape))
238
+ obj.thicknesses = numpy.array([], dtype=dtype)
239
+
240
+ state = ReconsState(
241
+ iter=IterState(0, 0, 0),
242
+ probe=probe,
243
+ object=obj,
244
+ scan=scan,
245
+ progress=ProgressState(iters=numpy.array([]), detector_errors=numpy.array([])),
246
+ wavelength=wavelength
247
+ )
248
+
249
+ # process post_init hooks
250
+ for p in plan.post_init:
251
+ (data, state) = p({
252
+ 'data': data, 'state': state,
253
+ 'dtype': dtype, 'seed': seed, 'xp': xp
254
+ })
255
+
256
+ # perform some checks on preprocessed data
257
+
258
+ if state.scan.shape[:-1] != data.patterns.shape[:-2]:
259
+ n_pos = int(numpy.prod(state.scan.shape[:-1]))
260
+ n_pat = int(numpy.prod(data.patterns.shape[:-2]))
261
+ if n_pos != n_pat:
262
+ raise ValueError(f"# of scan positions {n_pos} doesn't match # of patterns {n_pat}")
263
+
264
+ # reshape patterns to match scan
265
+ data.patterns = data.patterns.reshape((*state.scan.shape[:-1], *data.patterns.shape[-2:]))
266
+
267
+ avg_pattern_intensity = float(numpy.nanmean(numpy.nansum(data.patterns, axis=(-1, -2))))
268
+
269
+ if avg_pattern_intensity < 5.0:
270
+ logging.warning(
271
+ f"Mean pattern intensity is very low ({avg_pattern_intensity} particles). "
272
+ "Ensure that it is being scaled correctly to units of electrons/photons. "
273
+ "For simulated data, use the 'scale' or 'poisson' preprocessing"
274
+ )
275
+
276
+ return (data, state)
277
+
278
+
279
+ def prepare_for_engine(patterns: Patterns, state: ReconsState, xp: t.Any, engine: EnginePlan) -> t.Tuple[Patterns, ReconsState]:
280
+ # TODO: more graceful
281
+ if isinstance(engine, GradientEnginePlan) and not xp_is_jax(xp):
282
+ raise ValueError("The gradient descent engine requires the jax backend.")
283
+
284
+ state = state.to_xp(xp)
285
+
286
+ if engine.sim_shape is not None and engine.sim_shape != state.probe.data.shape[-2:]:
287
+ if engine.resize_method == 'pad_crop':
288
+ new_sampling = Sampling(engine.sim_shape, extent=tuple(state.probe.sampling.extent))
289
+ else:
290
+ new_sampling = Sampling(engine.sim_shape, sampling=tuple(state.probe.sampling.sampling))
291
+
292
+ logging.info(f"Resampling probe and patterns to shape {new_sampling.shape}...")
293
+ state.probe.data = state.probe.sampling.resample(state.probe.data, new_sampling)
294
+ # also resample patterns
295
+ patterns.patterns = state.probe.sampling.resample_recip(patterns.patterns, new_sampling)
296
+ # and pattern mask
297
+ patterns.pattern_mask = state.probe.sampling.resample_recip(patterns.pattern_mask, new_sampling)
298
+
299
+ state.probe.sampling = new_sampling
300
+
301
+ obj_sampling = state.object.sampling
302
+
303
+ if not numpy.allclose(state.probe.sampling.sampling, state.object.sampling.sampling):
304
+ # resample object -> probe
305
+ logging.info(f"Resampling object to pixel size {list(map(float, state.probe.sampling.sampling))}...")
306
+ obj_sampling = obj_sampling.with_sampling(state.probe.sampling.sampling)
307
+
308
+ obj_sampling_pad = obj_sampling.expand_to_scan(
309
+ state.scan, state.probe.sampling.extent / 2. + engine.obj_pad_px * state.probe.sampling.sampling
310
+ )
311
+
312
+ if obj_sampling_pad != obj_sampling:
313
+ logging.info(f"Padding object to shape {obj_sampling_pad.shape}")
314
+ obj_sampling = obj_sampling_pad
315
+
316
+ if obj_sampling != state.object.sampling:
317
+ state.object.data = state.object.sampling.resample(state.object.data, obj_sampling)
318
+ state.object.sampling = obj_sampling
319
+
320
+ current_probe_modes = state.probe.data.shape[0]
321
+ if engine.probe_modes != current_probe_modes:
322
+ # fix probe modes
323
+ if engine.probe_modes < current_probe_modes:
324
+ # TODO: redistribute intensity here
325
+ state.probe.data = state.probe.data[:engine.probe_modes]
326
+ else:
327
+ from phaser.utils.optics import make_hermetian_modes
328
+ if current_probe_modes != 1:
329
+ logging.info("Summing probe modes (in real-space) before recreating with different # of modes")
330
+
331
+ base_mode = xp.sum(state.probe.data, axis=0)
332
+ state.probe.data = make_hermetian_modes(base_mode, engine.probe_modes, base_mode_power=engine.base_mode_power)
333
+
334
+ if engine.slices is not None and (len(engine.slices.thicknesses) != len(state.object.thicknesses)
335
+ or not numpy.allclose(engine.slices.thicknesses, state.object.thicknesses)):
336
+ from phaser.utils.object import resample_slices
337
+ logging.info(f"Reslicing object from {max(1, len(state.object.thicknesses))} to {max(1, len(engine.slices.thicknesses))} slice(s)...")
338
+ state.object.data = resample_slices(state.object.data, state.object.thicknesses, engine.slices.thicknesses)
339
+ state.object.thicknesses = xp.array(engine.slices.thicknesses, dtype=state.object.thicknesses.dtype)
340
+
341
+ return patterns, state
342
+
343
+
344
+ _MISSING = object()
345
+
346
+
347
+ def merge(left: t.Any, right: t.Any) -> t.Any:
348
+ def _as_dict(val) -> t.Optional[dict]:
349
+ if isinstance(val, dict):
350
+ return val
351
+ if dataclasses.is_dataclass(val):
352
+ return dataclasses.asdict(val) # type: ignore
353
+ if isinstance(val, pane.PaneBase):
354
+ return val.dict(set_only=True)
355
+ return None
356
+
357
+ if left is _MISSING or right is _MISSING:
358
+ return left if right is _MISSING else right
359
+
360
+ if isinstance(left, Hook) and isinstance(right, Hook):
361
+ if left.ref != right.ref:
362
+ return right
363
+ d = merge(left.props or {}, right.props or {})
364
+ d['type'] = right.type if right.type is not None else right.ref
365
+ return pane.from_data(d, right.__class__)
366
+
367
+ if (left_d := _as_dict(left)) is not None and (right_d := _as_dict(right)) is not None:
368
+ keys = set(left_d.keys()) | set(right_d.keys())
369
+ return {k: merge(left_d.get(k, _MISSING), right_d.get(k, _MISSING)) for k in keys}
370
+
371
+ return left if right is _MISSING else right
@@ -0,0 +1,158 @@
1
+ from pathlib import Path
2
+ import typing as t
3
+
4
+ import numpy
5
+ from numpy.typing import NDArray, DTypeLike
6
+ import pane.annotations as annotations
7
+
8
+ from ..types import Dataclass, Slices
9
+ from .hook import Hook
10
+
11
+ if t.TYPE_CHECKING:
12
+ from phaser.utils.num import Sampling
13
+ from phaser.utils.object import ObjectSampling
14
+ from ..state import ObjectState, ProbeState, ReconsState, Patterns
15
+ from ..execute import Observer
16
+
17
+
18
+ class RawData(t.TypedDict):
19
+ patterns: NDArray[numpy.floating]
20
+ mask: NDArray[numpy.floating]
21
+ sampling: 'Sampling'
22
+ wavelength: t.Optional[float]
23
+ scan_hook: t.Union[t.Dict[str, t.Any], None]
24
+ probe_hook: t.Union[t.Dict[str, t.Any], None]
25
+ seed: t.Optional[object]
26
+
27
+
28
+ class LoadEmpadProps(Dataclass):
29
+ path: Path
30
+
31
+ diff_step: t.Optional[float] = None
32
+ kv: t.Optional[float] = None
33
+ adu: t.Optional[float] = None
34
+
35
+
36
+ class RawDataHook(Hook[None, RawData]):
37
+ known = {
38
+ 'empad': ('phaser.hooks.io.empad:load_empad', LoadEmpadProps),
39
+ }
40
+
41
+
42
+ class ProbeHookArgs(t.TypedDict):
43
+ sampling: 'Sampling'
44
+ wavelength: float
45
+ seed: t.Optional[object]
46
+ dtype: DTypeLike
47
+ xp: t.Any
48
+
49
+
50
+ class FocusedProbeProps(Dataclass):
51
+ defocus: t.Optional[float] = None # defocus, + is overfocus [A]
52
+ conv_angle: t.Optional[float] = None # semiconvergence angle [mrad]
53
+
54
+
55
+ class ProbeHook(Hook[ProbeHookArgs, 'ProbeState']):
56
+ known = {
57
+ 'focused': ('phaser.hooks.probe:focused_probe', FocusedProbeProps),
58
+ }
59
+
60
+
61
+ class ObjectHookArgs(t.TypedDict):
62
+ sampling: 'ObjectSampling'
63
+ wavelength: float
64
+ slices: t.Optional[Slices]
65
+ seed: t.Optional[object]
66
+ dtype: DTypeLike
67
+ xp: t.Any
68
+
69
+
70
+ class RandomObjectProps(Dataclass):
71
+ sigma: float = 1e-6
72
+
73
+
74
+ class ObjectHook(Hook[ObjectHookArgs, 'ObjectState']):
75
+ known = {
76
+ 'random': ('phaser.hooks.object:random_object', RandomObjectProps),
77
+ }
78
+
79
+
80
+ class ScanHookArgs(t.TypedDict):
81
+ seed: t.Optional[object]
82
+ dtype: DTypeLike
83
+ xp: t.Any
84
+
85
+
86
+ class RasterScanProps(Dataclass):
87
+ shape: t.Optional[t.Tuple[int, int]] = None # ny, nx (total shape)
88
+ step_size: t.Union[None, float, t.Tuple[float, float]] = None # A
89
+ rotation: t.Optional[float] = None # degrees CCW
90
+ affine: t.Optional[t.Annotated[NDArray[numpy.floating], annotations.shape((2, 2))]] = None
91
+
92
+
93
+ class ScanHook(Hook[ScanHookArgs, NDArray[numpy.floating]]):
94
+ known = {
95
+ 'raster': ('phaser.hooks.scan:raster_scan', RasterScanProps),
96
+ }
97
+
98
+
99
+ class PostInitArgs(t.TypedDict):
100
+ data: 'Patterns'
101
+ state: 'ReconsState'
102
+ seed: t.Optional[object]
103
+ dtype: DTypeLike
104
+ xp: t.Any
105
+
106
+
107
+ class ScaleProps(Dataclass):
108
+ scale: float
109
+
110
+
111
+ class CropDataProps(Dataclass):
112
+ crop: t.Tuple[
113
+ # y_i, y_f, x_i, x_f
114
+ t.Optional[int], t.Optional[int], t.Optional[int], t.Optional[int],
115
+ ]
116
+
117
+
118
+ class PoissonProps(Dataclass):
119
+ scale: t.Optional[float] = None
120
+ gaussian: t.Optional[float] = 1.0e-3
121
+
122
+
123
+ class DropNanProps(Dataclass):
124
+ threshold: float = 0.9
125
+
126
+
127
+ class DiffractionAlignProps(Dataclass):
128
+ ...
129
+
130
+
131
+ class PostLoadHook(Hook[RawData, RawData]):
132
+ known = {
133
+ 'crop_data': ('phaser.hooks.preprocessing:crop_data', CropDataProps),
134
+ 'poisson': ('phaser.hooks.preprocessing:add_poisson_noise', PoissonProps),
135
+ 'scale': ('phaser.hooks.preprocessing:scale_patterns', ScaleProps),
136
+ }
137
+
138
+
139
+ class PostInitHook(Hook[PostInitArgs, t.Tuple['Patterns', 'ReconsState']]):
140
+ known = {
141
+ 'drop_nans': ('phaser.hooks.preprocessing:drop_nan_patterns', DropNanProps),
142
+ 'diffraction_align': ('phaser.hooks.preprocessing:diffraction_align', DiffractionAlignProps),
143
+ }
144
+
145
+
146
+ class EngineArgs(t.TypedDict):
147
+ data: 'Patterns'
148
+ state: 'ReconsState'
149
+ dtype: DTypeLike
150
+ xp: t.Any
151
+ recons_name: str
152
+ engine_i: int
153
+ observer: 'Observer'
154
+ seed: t.Any
155
+
156
+
157
+ class EngineHook(Hook[EngineArgs, 'ReconsState']):
158
+ known = {} # filled in by plan.py
phaser/hooks/hook.py ADDED
@@ -0,0 +1,159 @@
1
+ from __future__ import annotations
2
+
3
+ import abc
4
+ import importlib
5
+ import typing as t
6
+
7
+ import pane
8
+ from pane.convert import ConverterHandlers, DataType
9
+ from pane.converters import Converter, make_converter
10
+ from pane.errors import ErrorNode, WrongTypeError, ParseInterrupt, ProductErrorNode
11
+
12
+ T = t.TypeVar('T')
13
+ U = t.TypeVar('U')
14
+
15
+ class Hook(t.Generic[T, U], abc.ABC):
16
+ known: t.ClassVar[t.Dict[str, t.Tuple[str, type]]] = {}
17
+
18
+ def __init__(
19
+ self, ref: str, props: t.Optional[t.Any] = None, type: t.Optional[str] = None,
20
+ ):
21
+ self.ref: str = ref
22
+ self.type: t.Optional[str] = type
23
+ self.f: t.Optional[t.Callable[..., U]] = None
24
+ self.props: t.Optional[t.Any] = props
25
+
26
+ def func_ref(self) -> str:
27
+ if self.type is not None:
28
+ return self.type
29
+ return self.ref
30
+
31
+ def _resolve_ref(self) -> t.Callable:
32
+ if ':' not in self.ref:
33
+ if self.ref in globals():
34
+ return globals()[self.ref]
35
+ raise ValueError(f"Can't resolve function reference '{self.ref}'.")
36
+
37
+ (module_path, func_name) = self.ref.split(':')
38
+ try:
39
+ module = importlib.import_module(module_path)
40
+ except ImportError as e:
41
+ e.add_note(f"While resolving function reference {self.ref}")
42
+ raise
43
+
44
+ try:
45
+ return getattr(module, func_name)
46
+ except AttributeError:
47
+ raise AttributeError(f"No function '{func_name}' found in module '{module_path}'")
48
+
49
+ def resolve(self) -> t.Callable[..., U]:
50
+ if self.f is None:
51
+ self.f = self._resolve_ref()
52
+ return self.f
53
+
54
+ def __call__(self, args: T) -> U:
55
+ return self.resolve()(args, props=self.props if self.props is not None else {})
56
+
57
+ def __getattr__(self, key: t.Any) -> t.Any:
58
+ if isinstance(self.props, dict):
59
+ try:
60
+ return self.props[key]
61
+ except KeyError:
62
+ raise AttributeError(name=key, obj=self.props)
63
+ return getattr(self.props, key)
64
+
65
+ def __repr__(self) -> str:
66
+ if self.props is not None:
67
+ return f"FuncRef({self.func_ref()!r}, {self.props!r})"
68
+ return f"FuncRef({self.func_ref()!r})"
69
+
70
+ @classmethod
71
+ def _converter(cls, *args: type, handlers: ConverterHandlers) -> HookConverter[T, U]:
72
+ return HookConverter(cls, handlers)
73
+
74
+
75
+ def _to_dict(val: t.Any) -> dict:
76
+ import dataclasses
77
+ import pane
78
+
79
+ if isinstance(val, dict):
80
+ return val
81
+ if dataclasses.is_dataclass(val):
82
+ return dataclasses.asdict(val) # type: ignore
83
+ if isinstance(val, pane.PaneBase):
84
+ return val.into_data() # type: ignore
85
+ return val.__dict__
86
+
87
+
88
+ class HookConverter(t.Generic[T, U], Converter[Hook[T, U]]):
89
+ def __init__(self, cls: t.Type[Hook[T, U]], handlers: ConverterHandlers):
90
+ self.cls = cls
91
+ self.inner: Converter[t.Union[str, t.Dict[str, t.Any]]] = make_converter(t.Union[str, t.Dict[str, t.Any]], handlers)
92
+
93
+ def expected(self, plural: bool = False) -> str:
94
+ if plural:
95
+ return "hooks to functions"
96
+ return "hook to function"
97
+
98
+ def into_data(self, val: Hook[T, U]) -> DataType:
99
+ if val.props is not None:
100
+ return {
101
+ 'type': val.func_ref(),
102
+ **_to_dict(val.props)
103
+ }
104
+ return val.func_ref()
105
+
106
+ def try_convert(self, val: t.Any) -> Hook[T, U]:
107
+ val = self.inner.try_convert(val)
108
+ if isinstance(val, str):
109
+ ref = val
110
+ props = {}
111
+ else:
112
+ if 'type' not in val:
113
+ raise ParseInterrupt()
114
+ ref = str(val.pop('type'))
115
+ props = val
116
+
117
+ if ref in self.cls.known:
118
+ ty = ref
119
+ (ref, props_ty) = self.cls.known[ty]
120
+
121
+ converter = make_converter(props_ty)
122
+ props = converter.try_convert(props)
123
+ elif ':' not in ref:
124
+ raise ParseInterrupt()
125
+ else:
126
+ ty = None
127
+
128
+ return self.cls(ref, props, ty)
129
+
130
+ def collect_errors(self, val: t.Any) -> t.Optional[ErrorNode]:
131
+ try:
132
+ val = self.inner.try_convert(val)
133
+ except ParseInterrupt:
134
+ return self.inner.collect_errors(val)
135
+ if isinstance(val, str):
136
+ ref = val
137
+ props = {}
138
+ else:
139
+ if 'type' not in val:
140
+ return ProductErrorNode(self.expected(), {}, val, set(['type']))
141
+ ref = str(val.pop('type'))
142
+ props = val
143
+
144
+ if ref in self.cls.known:
145
+ ty = ref
146
+ (ref, props_ty) = self.cls.known[ty]
147
+
148
+ converter = make_converter(props_ty)
149
+ try:
150
+ props = converter.try_convert(props)
151
+ except ParseInterrupt:
152
+ return converter.collect_errors(props)
153
+ elif ':' not in ref:
154
+ return WrongTypeError(
155
+ self.expected(), ref,
156
+ info=f"Known hooks: '{', '.join(self.cls.known.keys())}'"
157
+ )
158
+
159
+ return None