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
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
|
phaser/hooks/__init__.py
ADDED
|
@@ -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
|