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
|
@@ -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'])
|