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,451 @@
1
+ import contextlib
2
+ import logging
3
+ from functools import partial
4
+ from pathlib import Path
5
+ import typing as t
6
+
7
+ import numpy
8
+ from numpy.typing import NDArray
9
+ from typing_extensions import Self
10
+
11
+ from phaser.hooks.solver import NoiseModel
12
+ from phaser.utils.misc import create_sparse_groupings, create_compact_groupings, shuffled, jax_dataclass
13
+ from phaser.utils.num import (
14
+ get_array_module, cast_array_module, jit,
15
+ fft2, ifft2, abs2, check_finite, at, Float
16
+ )
17
+ from phaser.utils.optics import fourier_shift_filter
18
+ from phaser.utils.io import OutputDir
19
+ from phaser.execute import Observer
20
+ from phaser.state import ReconsState
21
+ from phaser.hooks import EngineArgs
22
+ from phaser.hooks.solver import GradientSolver
23
+ from phaser.hooks.regularization import CostRegularizer, GroupConstraint
24
+ from phaser.plan import GradientEnginePlan
25
+ from phaser.types import process_flag, flag_any_true, ReconsVar
26
+ from ..common.output import output_images, output_state
27
+ from ..common.simulation import GroupManager, stream_patterns, make_propagators, slice_forwards
28
+
29
+
30
+ logger = logging.getLogger(__name__)
31
+ _PER_ITER_VARS: t.FrozenSet[ReconsVar] = frozenset({'positions'})
32
+
33
+
34
+ def process_solvers(
35
+ plan: GradientEnginePlan
36
+ ) -> t.Tuple[t.FrozenSet[ReconsVar], t.Sequence[GradientSolver[t.Any]], t.Sequence[GradientSolver[t.Any]]]:
37
+ # process solvers, and split into per-group and per-iter solvers
38
+ solvers = plan.solvers
39
+
40
+ seen = set()
41
+ duplicate = set()
42
+
43
+ group_solvers = []
44
+ iter_solvers = []
45
+
46
+ for (vars, solver) in solvers.items():
47
+ if len(vars) == 0:
48
+ continue
49
+
50
+ duplicate |= vars & seen
51
+ seen |= vars
52
+
53
+ if vars <= _PER_ITER_VARS:
54
+ iter_solvers.append(solver({'plan': plan, 'params': vars}))
55
+ continue
56
+
57
+ if len(vars & _PER_ITER_VARS):
58
+ # TODO: is it easier to just split the solver here?
59
+ raise ValueError(f"The same solver can't handle both per-iteration "
60
+ f"({', '.join(map(repr, vars & _PER_ITER_VARS))}) and per-group "
61
+ f"({', '.join(map(repr, vars - _PER_ITER_VARS))}) variables")
62
+
63
+ group_solvers.append(solver({'plan': plan, 'params': vars}))
64
+
65
+ if len(duplicate):
66
+ raise ValueError(f"Duplicate solvers for variable(s) {', '.join(map(repr, duplicate))}.")
67
+
68
+ return (
69
+ frozenset(seen), tuple(group_solvers), tuple(iter_solvers)
70
+ )
71
+
72
+
73
+ _PATH_MAP: t.Dict[t.Tuple[str, ...], ReconsVar] = {
74
+ ('object', 'data'): 'object',
75
+ ('probe', 'data'): 'probe',
76
+ ('scan',): 'positions',
77
+ }
78
+
79
+ def extract_vars(state: ReconsState, vars: t.AbstractSet[ReconsVar], group: t.Optional[NDArray[numpy.integer]] = None) -> t.Tuple[t.Dict[ReconsVar, t.Any], ReconsState]:
80
+ import jax.tree_util
81
+
82
+ d = {}
83
+
84
+ def f(path: t.Tuple[str, ...], val: t.Any):
85
+ if (var := _PATH_MAP.get(path)) and var in vars:
86
+ if var in _PER_ITER_VARS and group is not None:
87
+ d[var] = val[tuple(group)]
88
+ else:
89
+ d[var] = val
90
+ return None
91
+ return val
92
+
93
+ state = jax.tree_util.tree_map_with_path(f, state, is_leaf=lambda x: x is None)
94
+ return (d, state)
95
+
96
+ def insert_vars(vars: t.Dict[ReconsVar, t.Any], state: ReconsState, group: t.Optional[NDArray[numpy.integer]] = None) -> ReconsState:
97
+ import jax.tree_util
98
+
99
+ def f(path: t.Tuple[str, ...], val: t.Any):
100
+ if (var := _PATH_MAP.get(path)):
101
+ if var in vars:
102
+ return vars[var]
103
+ if val is None:
104
+ raise ValueError(f"Missing value for var {var}")
105
+ if var in _PER_ITER_VARS and group is not None:
106
+ return val[tuple(group)]
107
+ return val
108
+
109
+ return jax.tree_util.tree_map_with_path(f, state, is_leaf=lambda x: x is None)
110
+
111
+
112
+ def apply_update(state: ReconsState, update: t.Dict[ReconsVar, numpy.ndarray]) -> ReconsState:
113
+ if 'probe' in update:
114
+ state.probe.data += update['probe']
115
+ if 'object' in update:
116
+ state.object.data += update['object']
117
+ if 'positions' in update:
118
+ xp = get_array_module(update['positions'])
119
+ # subtract mean position update
120
+ update['positions'] -= xp.mean(update['positions'], tuple(range(update['positions'].ndim - 1)))
121
+ logger.info(f"Position update: mean {xp.mean(xp.linalg.norm(update['positions'], axis=-1))}")
122
+ state.scan += update['positions']
123
+
124
+ return state
125
+
126
+
127
+ def filter_vars(d: t.Dict[ReconsVar, t.Any], vars: t.AbstractSet[ReconsVar]) -> t.Dict[ReconsVar, t.Any]:
128
+ return {k: v for (k, v) in d.items() if k in vars}
129
+
130
+
131
+ @jax_dataclass
132
+ class SolverStates:
133
+ noise_model_state: t.Any
134
+ group_solver_states: t.List[t.Any]
135
+ regularizer_states: t.List[t.Any]
136
+ group_constraint_states: t.List[t.Any]
137
+
138
+ @classmethod
139
+ def init_state(
140
+ cls, sim: ReconsState, xp: t.Any,
141
+ noise_model: NoiseModel,
142
+ group_solvers: t.Iterable[GradientSolver[t.Any]],
143
+ regularizers: t.Iterable[CostRegularizer[t.Any]],
144
+ group_constraints: t.Iterable[GroupConstraint[t.Any]],
145
+ ) -> Self:
146
+ noise_model_state = noise_model.init_state(sim)
147
+ group_solver_states = [solver.init_state(sim) for solver in group_solvers]
148
+ regularizer_states = [reg.init_state(sim) for reg in regularizers]
149
+ group_constraint_states = [reg.init_state(sim) for reg in group_constraints]
150
+
151
+ return cls(
152
+ noise_model_state, group_solver_states, regularizer_states, group_constraint_states
153
+ )
154
+
155
+
156
+ def run_engine(args: EngineArgs, props: GradientEnginePlan) -> ReconsState:
157
+ import jax
158
+ import jax.numpy
159
+ from optax.tree_utils import tree_zeros_like
160
+ jax.config.update('jax_traceback_filtering', 'off')
161
+
162
+ xp = cast_array_module(jax.numpy)
163
+ dtype = t.cast(t.Type[numpy.floating], args['dtype'])
164
+
165
+ observer: Observer = args.get('observer', [])
166
+ recons_name = args['recons_name']
167
+ engine_i = args['engine_i']
168
+
169
+ logger.info(f"Starting engine #{args['engine_i'] + 1}...")
170
+
171
+ state = args['state']
172
+ seed = args['seed']
173
+ patterns = args['data'].patterns
174
+ pattern_mask = args['data'].pattern_mask
175
+
176
+ noise_model = props.noise_model(None)
177
+
178
+ (all_vars, group_solvers, iter_solvers) = process_solvers(props)
179
+
180
+ regularizers = tuple(reg(None) for reg in props.regularizers)
181
+ group_constraints = tuple(reg(None) for reg in props.group_constraints)
182
+ iter_constraints = tuple(reg(None) for reg in props.iter_constraints)
183
+
184
+ flags = {
185
+ 'probe': process_flag(props.update_probe),
186
+ 'object': process_flag(props.update_object),
187
+ 'positions': process_flag(props.update_positions),
188
+ }
189
+ save = process_flag(props.save)
190
+ save_images = process_flag(props.save_images)
191
+ # shuffle_groups defaults to True for sparse groups, False for compact groups
192
+ shuffle_groups = process_flag(props.shuffle_groups or not props.compact)
193
+ groups = GroupManager(state.scan, props.grouping, props.compact, seed)
194
+
195
+ any_output = flag_any_true(save, props.niter) or flag_any_true(save_images, props.niter)
196
+
197
+ # TODO: this really needs cleanup
198
+
199
+
200
+ with OutputDir(
201
+ props.save_options.out_dir, any_output,
202
+ engine_i=engine_i, name=recons_name,
203
+ group=groups.grouping, niter=props.niter,
204
+ noise_model=noise_model.name(),
205
+ ) as out_dir:
206
+ propagators = make_propagators(state, props.bwlim_frac)
207
+
208
+ start_i = state.iter.total_iter
209
+ observer.init_solver(state, engine_i)
210
+
211
+ # runs rescaling
212
+ rescale_factors = []
213
+ for (group_i, group) in enumerate(groups.iter(state.scan)):
214
+ group_rescale_factors = dry_run(state, group, propagators, patterns, xp=xp, dtype=dtype)
215
+ rescale_factors.append(group_rescale_factors)
216
+
217
+ rescale_factors = xp.concatenate(rescale_factors, axis=0)
218
+ rescale_factor = xp.mean(rescale_factors)
219
+
220
+ logger.info("Pre-calculated intensities")
221
+ logger.info(f"Rescaling initial probe intensity by {rescale_factor:.2e}")
222
+ state.probe.data *= xp.sqrt(rescale_factor)
223
+ probe_int = xp.sum(abs2(state.probe.data))
224
+
225
+ observer.start_solver()
226
+
227
+ solver_states = SolverStates.init_state(state, xp, noise_model, group_solvers, regularizers, group_constraints)
228
+ iter_solver_states = [solver.init_state(state) for solver in iter_solvers]
229
+ iter_constraint_states = [reg.init_state(state) for reg in iter_constraints]
230
+
231
+ #with jax.profiler.trace("/tmp/jax-trace", create_perfetto_link=True):
232
+ for i in range(1, props.niter+1):
233
+ losses = []
234
+
235
+ # mask vars we're updating this iteration
236
+ iter_vars = all_vars & t.cast(t.Set[ReconsVar],
237
+ set(k for (k, flag) in flags.items() if flag({'state': state, 'niter': props.niter}))
238
+ )
239
+ # gradients for per-iteration solvers
240
+ iter_grads = tree_zeros_like(extract_vars(state, iter_vars & _PER_ITER_VARS)[0])
241
+ # whether to shuffle groups this iteration
242
+ iter_shuffle_groups = shuffle_groups({'state': state, 'niter': props.niter})
243
+
244
+ # update schedules for this iteration
245
+ # this needs to be done outside the JIT context, which makes this kinda hacky
246
+ solver_states.group_solver_states = [
247
+ solver.update_for_iter(state, solver_state, props.niter)
248
+ for (solver, solver_state) in zip(group_solvers, solver_states.group_solver_states)
249
+ ]
250
+ iter_solver_states = [
251
+ solver.update_for_iter(state, solver_state, props.niter)
252
+ for (solver, solver_state) in zip(iter_solvers, iter_solver_states)
253
+ ]
254
+
255
+ for (group_i, (group, group_patterns)) in enumerate(stream_patterns(groups.iter(state.scan, i, iter_shuffle_groups),
256
+ patterns, xp=xp, buf_n=props.buffer_n_groups)):
257
+ (state, loss, iter_grads, solver_states) = run_group(
258
+ state, group=group, vars=iter_vars,
259
+ noise_model=noise_model,
260
+ group_solvers=group_solvers,
261
+ group_constraints=group_constraints,
262
+ regularizers=regularizers,
263
+ iter_grads=iter_grads,
264
+ solver_states=solver_states,
265
+ props=propagators,
266
+ group_patterns=group_patterns, #load_group(group),
267
+ pattern_mask=pattern_mask,
268
+ probe_int=probe_int,
269
+ xp=xp, dtype=dtype
270
+ )
271
+
272
+ losses.append(loss)
273
+ check_finite(state.object.data, state.probe.data, context=f"object or probe, group {group_i}")
274
+ observer.update_group(state, props.send_every_group)
275
+
276
+ loss = float(numpy.mean(losses))
277
+
278
+ # update per-iteration solvers
279
+ for (sol_i, solver) in enumerate(iter_solvers):
280
+ solver_grads = filter_vars(iter_grads, solver.params)
281
+ if len(solver_grads) == 0:
282
+ continue
283
+ (update, iter_solver_states[sol_i]) = solver.update(
284
+ state, iter_solver_states[sol_i], filter_vars(iter_grads, solver.params), loss
285
+ )
286
+ state = apply_update(state, update)
287
+
288
+ for (reg_i, reg) in enumerate(iter_constraints):
289
+ (state, iter_constraint_states[reg_i]) = reg.apply_iter(
290
+ state, iter_constraint_states[reg_i]
291
+ )
292
+
293
+ if 'positions' in iter_vars:
294
+ # check positions are at least overlapping object
295
+ state.object.sampling.check_scan(state.scan, state.probe.sampling.extent / 2.)
296
+
297
+ observer.update_iteration(state, i, props.niter, loss)
298
+
299
+ state.progress.iters = numpy.concatenate([state.progress.iters, [i + start_i]])
300
+ state.progress.detector_errors = numpy.concatenate([state.progress.detector_errors, [loss]])
301
+
302
+ if save({'state': state, 'niter': props.niter}):
303
+ output_state(state, out_dir, props.save_options)
304
+
305
+ if save_images({'state': state, 'niter': props.niter}):
306
+ output_images(state, out_dir, props.save_options)
307
+
308
+ observer.finish_solver()
309
+ return state
310
+
311
+
312
+ @partial(
313
+ jit,
314
+ static_argnames=('vars', 'xp', 'dtype', 'noise_model', 'group_solvers', 'group_constraints', 'regularizers'),
315
+ donate_argnames=('state', 'iter_grads', 'solver_states'),
316
+ )
317
+ def run_group(
318
+ state: ReconsState,
319
+ group: NDArray[numpy.integer],
320
+ vars: t.AbstractSet[ReconsVar], *,
321
+ noise_model: NoiseModel[t.Any],
322
+ group_solvers: t.Sequence[GradientSolver[t.Any]],
323
+ group_constraints: t.Sequence[GroupConstraint[t.Any]],
324
+ regularizers: t.Sequence[CostRegularizer[t.Any]],
325
+ iter_grads: t.Dict[ReconsVar, t.Any],
326
+ solver_states: SolverStates,
327
+ props: t.Optional[NDArray[numpy.complexfloating]],
328
+ group_patterns: NDArray[numpy.floating],
329
+ pattern_mask: NDArray[numpy.floating],
330
+ probe_int: t.Union[float, numpy.floating],
331
+ xp: t.Any,
332
+ dtype: t.Type[numpy.floating],
333
+ ) -> t.Tuple[ReconsState, float, t.Dict[ReconsVar, t.Any], SolverStates]:
334
+ import jax
335
+ xp = cast_array_module(xp)
336
+
337
+ ((loss, solver_states), grad) = jax.value_and_grad(run_model, has_aux=True)(
338
+ *extract_vars(state, vars, group),
339
+ group=group, props=props, group_patterns=group_patterns, pattern_mask=pattern_mask,
340
+ noise_model=noise_model, regularizers=regularizers, solver_states=solver_states,
341
+ xp=xp, dtype=dtype
342
+ )
343
+ # steepest descent direction
344
+ grad = jax.tree.map(lambda v: -v.conj(), grad, is_leaf=lambda x: x is None)
345
+ for k in grad.keys():
346
+ if k == 'probe':
347
+ grad[k] /= group.shape[-1]
348
+ else:
349
+ grad[k] /= probe_int * group.shape[-1]
350
+
351
+ # update iter grads at group
352
+ iter_grads = jax.tree.map(lambda v1, v2: at(v1, tuple(group)).set(v2), iter_grads, filter_vars(grad, vars & _PER_ITER_VARS))
353
+
354
+ for (sol_i, solver) in enumerate(group_solvers):
355
+ solver_grads = filter_vars(grad, solver.params)
356
+ if len(solver_grads) == 0:
357
+ continue
358
+ (update, solver_states.group_solver_states[sol_i]) = solver.update(
359
+ state, solver_states.group_solver_states[sol_i], solver_grads, loss
360
+ )
361
+ state = apply_update(state, update)
362
+
363
+ for (reg_i, reg) in enumerate(group_constraints):
364
+ (state, solver_states.group_constraint_states[reg_i]) = reg.apply_group(
365
+ group, state, solver_states.group_constraint_states[reg_i]
366
+ )
367
+
368
+ return (state, loss, iter_grads, solver_states)
369
+
370
+
371
+ @partial(
372
+ jit,
373
+ static_argnames=('xp', 'dtype', 'noise_model', 'regularizers'),
374
+ donate_argnames=('solver_states',),
375
+ )
376
+ def run_model(
377
+ vars: t.Dict[ReconsVar, t.Any],
378
+ sim: ReconsState,
379
+ group: NDArray[numpy.integer],
380
+ props: t.Optional[NDArray[numpy.complexfloating]],
381
+ group_patterns: NDArray[numpy.floating],
382
+ pattern_mask: NDArray[numpy.floating],
383
+ noise_model: NoiseModel[t.Any],
384
+ regularizers: t.Sequence[CostRegularizer[t.Any]],
385
+ solver_states: SolverStates,
386
+ xp: t.Any,
387
+ dtype: t.Type[numpy.floating],
388
+ ) -> t.Tuple[Float, SolverStates]:
389
+ # apply vars to simulation
390
+ sim = insert_vars(vars, sim, group)
391
+ group_scan = sim.scan
392
+
393
+ (ky, kx) = sim.probe.sampling.recip_grid(dtype=dtype, xp=xp)
394
+ probes = sim.probe.data
395
+ group_obj = sim.object.sampling.get_view_at_pos(sim.object.data, group_scan, probes.shape[-2:])
396
+ group_subpx_filters = fourier_shift_filter(ky, kx, sim.object.sampling.get_subpx_shifts(group_scan, probes.shape[-2:]))[:, None, ...]
397
+ probes = ifft2(fft2(probes) * group_subpx_filters)
398
+
399
+ def sim_slice(slice_i: int, prop: t.Optional[NDArray[numpy.complexfloating]], psi):
400
+ if prop is not None:
401
+ return ifft2(fft2(psi * group_obj[:, slice_i, None]) * prop)
402
+
403
+ return psi * group_obj[:, slice_i, None]
404
+
405
+ model_wave = fft2(slice_forwards(props, probes, sim_slice))
406
+ model_intensity = xp.sum(abs2(model_wave), axis=1)
407
+
408
+ (loss, solver_states.noise_model_state) = noise_model.calc_loss(
409
+ model_wave, model_intensity, group_patterns, pattern_mask, solver_states.noise_model_state
410
+ )
411
+
412
+ for (reg_i, reg) in enumerate(regularizers):
413
+ (reg_loss, solver_states.regularizer_states[reg_i]) = reg.calc_loss_group(
414
+ group, sim, solver_states.regularizer_states[reg_i]
415
+ )
416
+ loss += reg_loss
417
+
418
+ return (loss, solver_states)
419
+
420
+
421
+ # TODO: DRY
422
+ @partial(
423
+ jit,
424
+ static_argnames=('xp', 'dtype'),
425
+ )
426
+ def dry_run(
427
+ sim: ReconsState,
428
+ group: NDArray[numpy.integer],
429
+ props: t.Optional[NDArray[numpy.complexfloating]],
430
+ patterns: NDArray[numpy.floating],
431
+ xp: t.Any,
432
+ dtype: t.Type[numpy.floating],
433
+ ) -> NDArray[numpy.floating]:
434
+ (ky, kx) = sim.probe.sampling.recip_grid(dtype=dtype, xp=xp)
435
+ probes = sim.probe.data
436
+ group_obj = sim.object.sampling.get_view_at_pos(sim.object.data, sim.scan[tuple(group)], probes.shape[-2:])
437
+ group_subpx_filters = fourier_shift_filter(ky, kx, sim.object.sampling.get_subpx_shifts(sim.scan[tuple(group)], probes.shape[-2:]))[:, None, ...]
438
+ probes = ifft2(fft2(probes) * group_subpx_filters)
439
+
440
+ def sim_slice(slice_i: int, prop: t.Optional[NDArray[numpy.complexfloating]], psi):
441
+ if prop is not None:
442
+ return ifft2(fft2(psi * group_obj[:, slice_i, None]) * prop)
443
+
444
+ return psi * group_obj[:, slice_i, None]
445
+
446
+ model_wave = fft2(slice_forwards(props, probes, sim_slice))
447
+ # sum over incoherent modes and over the pattern
448
+ model_intensity = xp.sum(abs2(model_wave), axis=(1, -2, -1))
449
+ exp_intensity = xp.sum(xp.array(patterns[tuple(group)]), axis=(-2, -1))
450
+
451
+ return exp_intensity / model_intensity
@@ -0,0 +1,139 @@
1
+ import logging
2
+ import typing as t
3
+
4
+ import numpy
5
+ from numpy.typing import NDArray, ArrayLike
6
+
7
+ from phaser.utils.num import as_array, abs2
8
+ from phaser.hooks.solver import GradientSolver, GradientSolverArgs
9
+ from phaser.hooks.schedule import FlagArgs, ScheduleLike
10
+ from phaser.types import ReconsVar, process_schedule
11
+ from phaser.plan import GradientEnginePlan, AdamSolverPlan, PolyakSGDSolverPlan, SGDSolverPlan
12
+ from phaser.state import ReconsState
13
+ from .run import extract_vars, apply_update
14
+
15
+ import optax
16
+ from optax import GradientTransformation, GradientTransformationExtraArgs
17
+ from optax.schedules import StatefulSchedule
18
+
19
+
20
+ class OptaxScheduleWrapper(StatefulSchedule):
21
+ def __init__(self, schedule: ScheduleLike):
22
+ self.inner = process_schedule(schedule)
23
+
24
+ def init(self) -> t.Optional[float]:
25
+ return None
26
+
27
+ def update_for_iter(self, sim: ReconsState, state: t.Optional[float], niter: int) -> float:
28
+ return self.inner({'state': sim, 'niter': niter})
29
+
30
+ # mock update from inside jax
31
+ def update(
32
+ self, state: t.Optional[float],
33
+ **extra_args,
34
+ ) -> t.Optional[float]:
35
+ return state
36
+
37
+ def __call__(
38
+ self, state: t.Optional[float],
39
+ **extra_args,
40
+ ) -> float:
41
+ assert state is not None
42
+ return state
43
+
44
+
45
+ OptaxSolverState: t.TypeAlias = t.Tuple[t.Any, t.Dict[str, t.Optional[float]]]
46
+
47
+
48
+ class OptaxSolver(GradientSolver[OptaxSolverState]):
49
+ def __init__(self, name: str, factory: t.Callable[..., GradientTransformation], hyperparams: t.Mapping[str, ScheduleLike],
50
+ params: t.Iterable[ReconsVar]):
51
+ self.factory: t.Callable[..., GradientTransformation] = factory
52
+ #self.inner: GradientTransformationExtraArgs = optax.with_extra_args_support(solver)
53
+
54
+ self.hyperparams: t.Dict[str, OptaxScheduleWrapper] = {k: OptaxScheduleWrapper(v) for (k, v) in hyperparams.items()}
55
+ self.params: t.FrozenSet[ReconsVar] = frozenset(params)
56
+
57
+ self.name: str = name # or self.inner.__class__.__name__
58
+
59
+ def init_state(self, sim: ReconsState) -> OptaxSolverState:
60
+ return (
61
+ None,
62
+ {k: v.init() for (k, v) in self.hyperparams.items()},
63
+ )
64
+
65
+ def _resolve(self, hparams: t.Mapping[str, t.Optional[float]]) -> GradientTransformationExtraArgs:
66
+ return optax.with_extra_args_support(
67
+ self.factory(**{k: v(hparams[k]) for (k, v) in self.hyperparams.items()})
68
+ )
69
+
70
+ def update_for_iter(self, sim: ReconsState, state: OptaxSolverState, niter: int) -> OptaxSolverState:
71
+ hparams_state: t.Dict[str, t.Optional[float]] = {k: v.update_for_iter(sim, state[1][k], niter) for (k, v) in self.hyperparams.items()}
72
+ return (
73
+ self._resolve(hparams_state).init(params=extract_vars(sim, self.params)[0]) if state[0] is None else state[0],
74
+ hparams_state
75
+ )
76
+
77
+ def update(
78
+ self, sim: 'ReconsState', state: OptaxSolverState, grad: t.Dict[ReconsVar, numpy.ndarray], loss: float,
79
+ ) -> t.Tuple[t.Dict[ReconsVar, numpy.ndarray], OptaxSolverState]:
80
+ (inner_state, hparams_state) = state
81
+ hparams_state = {k: v.update(hparams_state[k]) for (k, v) in self.hyperparams.items()}
82
+ (updates, inner_state) = self._resolve(hparams_state).update(
83
+ grad, inner_state, params=extract_vars(sim, self.params)[0], value=loss, loss=loss
84
+ )
85
+ return (t.cast(t.Dict[ReconsVar, t.Any], updates), (inner_state, hparams_state))
86
+
87
+
88
+ class SGDSolver(OptaxSolver):
89
+ def __init__(self, args: GradientSolverArgs, props: SGDSolverPlan):
90
+ hparams = {
91
+ 'learning_rate': props.learning_rate
92
+ }
93
+
94
+ if props.momentum is not None:
95
+ hparams['momentum'] = props.momentum
96
+ def factory(**kwargs: t.Any) -> GradientTransformation:
97
+ return optax.chain(
98
+ optax.trace(kwargs['momentum'], props.nesterov),
99
+ optax.scale_by_learning_rate(kwargs['learning_rate'], flip_sign=False),
100
+ )
101
+ else:
102
+ def factory(**kwargs: t.Any) -> GradientTransformation:
103
+ return optax.scale_by_learning_rate(kwargs['learning_rate'], flip_sign=False)
104
+
105
+ super().__init__('sgd', factory, hparams, args['params'])
106
+
107
+
108
+ class AdamSolver(OptaxSolver):
109
+ def __init__(self, args: GradientSolverArgs, props: AdamSolverPlan):
110
+ hparams = {
111
+ 'learning_rate': props.learning_rate
112
+ }
113
+
114
+ def factory(**kwargs) -> GradientTransformation:
115
+ return optax.chain(
116
+ optax.scale_by_adam(props.b1, props.b2, props.eps, props.eps_root, nesterov=props.nesterov),
117
+ optax.scale_by_learning_rate(learning_rate=kwargs['learning_rate'], flip_sign=False),
118
+ )
119
+
120
+ super().__init__('adam', factory, hparams, args['params'])
121
+
122
+
123
+ class PolyakSGDSolver(OptaxSolver):
124
+ def __init__(self, args: GradientSolverArgs, props: PolyakSGDSolverPlan):
125
+ hparams = {
126
+ 'max_learning_rate': props.max_learning_rate,
127
+ 'scaling': props.scaling,
128
+ }
129
+
130
+ def factory(**kwargs) -> GradientTransformation:
131
+ return optax.chain(
132
+ optax.scale_by_learning_rate(kwargs['scaling'], flip_sign=False),
133
+ optax.scale_by_polyak(
134
+ max_learning_rate=kwargs['max_learning_rate'], f_min=props.f_min,
135
+ eps=props.eps, #variant='sps',
136
+ )
137
+ )
138
+
139
+ super().__init__('polyak_sgd', factory, hparams, args['params'])