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