solvephase 0.1.0__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.
- solvephase/__about__.py +3 -0
- solvephase/__init__.py +110 -0
- solvephase/algorithms/__init__.py +4 -0
- solvephase/algorithms/cdi.py +967 -0
- solvephase/algorithms/fast_furious.py +494 -0
- solvephase/algorithms/gerchberg_saxton.py +255 -0
- solvephase/algorithms/lift.py +553 -0
- solvephase/algorithms/phase_diversity.py +658 -0
- solvephase/algorithms/tie.py +693 -0
- solvephase/algorithms/wirtinger.py +865 -0
- solvephase/api.py +228 -0
- solvephase/backend.py +24 -0
- solvephase/basis.py +516 -0
- solvephase/cli.py +131 -0
- solvephase/focal.py +367 -0
- solvephase/interop.py +173 -0
- solvephase/losses.py +169 -0
- solvephase/metrics.py +5 -0
- solvephase/operators.py +391 -0
- solvephase/optimize.py +404 -0
- solvephase/propagation.py +26 -0
- solvephase/pupil.py +5 -0
- solvephase/py.typed +0 -0
- solvephase/result.py +228 -0
- solvephase/retrieval.py +627 -0
- solvephase/simulate.py +113 -0
- solvephase/unwrap.py +5 -0
- solvephase-0.1.0.dist-info/METADATA +198 -0
- solvephase-0.1.0.dist-info/RECORD +32 -0
- solvephase-0.1.0.dist-info/WHEEL +4 -0
- solvephase-0.1.0.dist-info/entry_points.txt +2 -0
- solvephase-0.1.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,967 @@
|
|
|
1
|
+
"""Coherent diffraction imaging (CDI): iterative projection algorithms.
|
|
2
|
+
|
|
3
|
+
A compact object is illuminated coherently and only the modulus of its
|
|
4
|
+
far-field (Fourier) transform is measured. When the diffraction pattern is
|
|
5
|
+
oversampled (the object occupies less than half of the reconstruction window
|
|
6
|
+
along each axis, Miao, Sayre & Chapman, JOSA A 15, 1662, 1998) the phase is
|
|
7
|
+
fixed by the measured modulus together with an object-domain constraint (a
|
|
8
|
+
support, and optionally realness or positivity). The algorithms here
|
|
9
|
+
alternate projections onto the two constraint sets:
|
|
10
|
+
|
|
11
|
+
* ``P_M`` - the Fourier-modulus projection: keep the phase of the current
|
|
12
|
+
transform and impose the measured modulus on measured pixels;
|
|
13
|
+
* ``P_S`` - the object-domain projection: zero outside the support and, inside
|
|
14
|
+
it, the nearest complex/real/non-negative value (and amplitude bounds).
|
|
15
|
+
|
|
16
|
+
Implemented update rules (``x`` is the iterate, ``y = P_M x``,
|
|
17
|
+
``R = 2 P - I``):
|
|
18
|
+
|
|
19
|
+
* **ER**, error reduction ``x <- P_S P_M x`` (Gerchberg & Saxton, Optik 35,
|
|
20
|
+
237, 1972; Fienup, Opt. Lett. 3, 27, 1978; Appl. Opt. 21, 2758, 1982);
|
|
21
|
+
* **HIO**, hybrid input-output (Fienup 1982): ``y`` where it satisfies the
|
|
22
|
+
object constraints, ``x - beta y`` elsewhere;
|
|
23
|
+
* **DM**, the difference map (Elser, JOSA A 20, 40, 2003) with
|
|
24
|
+
``gamma_S = -1/beta`` and ``gamma_M = 1/beta``;
|
|
25
|
+
* **RAAR**, relaxed averaged alternating reflections (Luke, Inverse Problems
|
|
26
|
+
21, 37, 2005);
|
|
27
|
+
* **RRR**, relax-reflect-reflect (Elser, Lan & Bendory, SIAM J. Imaging
|
|
28
|
+
Sci. 11, 2429, 2018);
|
|
29
|
+
* **ASR**, averaged successive reflections = Douglas-Rachford (Bauschke,
|
|
30
|
+
Combettes & Luke, JOSA A 19, 1334, 2002);
|
|
31
|
+
* **HPR**, hybrid projection-reflection (Bauschke, Combettes & Luke, JOSA A
|
|
32
|
+
20, 1025, 2003);
|
|
33
|
+
* **OSS**, oversampling smoothness (Rodriguez, Xu, Chen, Zou & Miao, J.
|
|
34
|
+
Appl. Cryst. 46, 312, 2013): HIO whose out-of-support region is low-pass
|
|
35
|
+
filtered by a Gaussian of decreasing width;
|
|
36
|
+
* **shrinkwrap** support refinement (Marchesini et al., Phys. Rev. B 68,
|
|
37
|
+
140101(R), 2003), usable with any of the above.
|
|
38
|
+
|
|
39
|
+
Algorithms are chained with a schedule (``"hio:400,er:100"``), and many
|
|
40
|
+
random starts run simultaneously as a leading batch dimension, so each
|
|
41
|
+
iteration is one batched FFT pair. The far-field arrays are centred (DC at
|
|
42
|
+
``n // 2``); they are shifted once on entry and once on exit, never inside
|
|
43
|
+
the loop. Sub-pixel registration in :func:`align_object` follows
|
|
44
|
+
Guizar-Sicairos, Thurman & Fienup, Opt. Lett. 33, 156 (2008).
|
|
45
|
+
"""
|
|
46
|
+
|
|
47
|
+
from __future__ import annotations
|
|
48
|
+
|
|
49
|
+
import math
|
|
50
|
+
import time
|
|
51
|
+
from collections.abc import Callable, Sequence
|
|
52
|
+
from dataclasses import dataclass, fields, replace
|
|
53
|
+
from typing import Any
|
|
54
|
+
|
|
55
|
+
import numpy as np
|
|
56
|
+
|
|
57
|
+
from ..backend import Backend, BackendLike, backend_of, get_backend, to_numpy
|
|
58
|
+
|
|
59
|
+
__all__ = [
|
|
60
|
+
"CDI_ALGORITHMS",
|
|
61
|
+
"DEFAULT_BETA",
|
|
62
|
+
"CDIResult",
|
|
63
|
+
"ShrinkwrapConfig",
|
|
64
|
+
"SimulatedCDI",
|
|
65
|
+
"align_object",
|
|
66
|
+
"autocorrelation_support",
|
|
67
|
+
"cdi",
|
|
68
|
+
"parse_schedule",
|
|
69
|
+
"simulate_cdi",
|
|
70
|
+
]
|
|
71
|
+
|
|
72
|
+
CDI_ALGORITHMS = ("er", "hio", "dm", "raar", "rrr", "asr", "hpr", "oss")
|
|
73
|
+
"""Names accepted in a :func:`cdi` schedule."""
|
|
74
|
+
|
|
75
|
+
_ALIASES = {
|
|
76
|
+
"douglas-rachford": "asr",
|
|
77
|
+
"dr": "asr",
|
|
78
|
+
"difference-map": "dm",
|
|
79
|
+
"difference_map": "dm",
|
|
80
|
+
"error-reduction": "er",
|
|
81
|
+
}
|
|
82
|
+
_CONSTRAINTS = ("complex", "real", "positive")
|
|
83
|
+
DEFAULT_BETA = {
|
|
84
|
+
"er": 1.0,
|
|
85
|
+
"hio": 0.9,
|
|
86
|
+
"dm": 0.9,
|
|
87
|
+
"raar": 0.98,
|
|
88
|
+
"rrr": 0.5,
|
|
89
|
+
"asr": 1.0,
|
|
90
|
+
"hpr": 0.9,
|
|
91
|
+
"oss": 0.9,
|
|
92
|
+
}
|
|
93
|
+
"""Per-algorithm ``beta`` used when neither the stage nor :func:`cdi` sets one.
|
|
94
|
+
|
|
95
|
+
ER and ASR have no parameter (their entry is informational)."""
|
|
96
|
+
|
|
97
|
+
ScheduleLike = str | Sequence[Any]
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
# ----------------------------------------------------------------------------- configuration
|
|
101
|
+
@dataclass(frozen=True)
|
|
102
|
+
class ShrinkwrapConfig:
|
|
103
|
+
"""Shrinkwrap support refinement (Marchesini et al. 2003).
|
|
104
|
+
|
|
105
|
+
Every :attr:`every` iterations the support becomes
|
|
106
|
+
``G_sigma * |object| > threshold * max(G_sigma * |object|)``, where
|
|
107
|
+
``G_sigma`` is a Gaussian of standard deviation ``sigma`` pixels and
|
|
108
|
+
``|object|`` is the modulus of the current data-consistent image
|
|
109
|
+
``P_M x``. ``sigma`` starts at :attr:`sigma` and is multiplied by
|
|
110
|
+
:attr:`decay` after each update until it reaches :attr:`sigma_min`.
|
|
111
|
+
|
|
112
|
+
Attributes
|
|
113
|
+
----------
|
|
114
|
+
every:
|
|
115
|
+
Iterations between support updates.
|
|
116
|
+
threshold:
|
|
117
|
+
Fraction of the maximum of the blurred modulus kept in the support.
|
|
118
|
+
sigma:
|
|
119
|
+
Initial Gaussian standard deviation in pixels.
|
|
120
|
+
sigma_min:
|
|
121
|
+
Smallest standard deviation in pixels.
|
|
122
|
+
decay:
|
|
123
|
+
Factor applied to ``sigma`` after every update (``<= 1``).
|
|
124
|
+
start:
|
|
125
|
+
No update happens before this (global) iteration.
|
|
126
|
+
stop:
|
|
127
|
+
No update happens after this iteration (``None``: never stop). Useful
|
|
128
|
+
to let a final ER stage refine with a fixed support.
|
|
129
|
+
"""
|
|
130
|
+
|
|
131
|
+
every: int = 20
|
|
132
|
+
threshold: float = 0.2
|
|
133
|
+
sigma: float = 3.0
|
|
134
|
+
sigma_min: float = 1.5
|
|
135
|
+
decay: float = 0.99
|
|
136
|
+
start: int = 0
|
|
137
|
+
stop: int | None = None
|
|
138
|
+
|
|
139
|
+
def __post_init__(self) -> None:
|
|
140
|
+
if self.every < 1:
|
|
141
|
+
raise ValueError(f"shrinkwrap every must be >= 1, got {self.every}")
|
|
142
|
+
if not 0.0 < self.threshold < 1.0:
|
|
143
|
+
raise ValueError(f"shrinkwrap threshold must be in (0, 1), got {self.threshold}")
|
|
144
|
+
if self.sigma < 0 or self.sigma_min < 0 or self.sigma_min > self.sigma:
|
|
145
|
+
raise ValueError("shrinkwrap needs 0 <= sigma_min <= sigma (pixels)")
|
|
146
|
+
if not 0.0 < self.decay <= 1.0:
|
|
147
|
+
raise ValueError(f"shrinkwrap decay must be in (0, 1], got {self.decay}")
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
ShrinkwrapLike = ShrinkwrapConfig | dict | bool | None
|
|
151
|
+
|
|
152
|
+
|
|
153
|
+
def _shrinkwrap_config(value: ShrinkwrapLike) -> ShrinkwrapConfig | None:
|
|
154
|
+
if value is None or value is False:
|
|
155
|
+
return None
|
|
156
|
+
if value is True:
|
|
157
|
+
return ShrinkwrapConfig()
|
|
158
|
+
if isinstance(value, ShrinkwrapConfig):
|
|
159
|
+
return value
|
|
160
|
+
if isinstance(value, dict):
|
|
161
|
+
return ShrinkwrapConfig(**value)
|
|
162
|
+
raise TypeError(
|
|
163
|
+
f"shrinkwrap must be None, True, a dict or a ShrinkwrapConfig, got {type(value).__name__}"
|
|
164
|
+
)
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
def parse_schedule(
|
|
168
|
+
schedule: ScheduleLike, beta: float | None = None
|
|
169
|
+
) -> list[tuple[str, int, float]]:
|
|
170
|
+
"""Normalize a CDI schedule to ``[(algorithm, iterations, beta), ...]``.
|
|
171
|
+
|
|
172
|
+
Parameters
|
|
173
|
+
----------
|
|
174
|
+
schedule:
|
|
175
|
+
Either a string ``"hio:400,er:100"`` (an optional third field sets the
|
|
176
|
+
stage's beta: ``"raar:300:0.8"``) or a sequence of ``(name, n)`` /
|
|
177
|
+
``(name, n, beta)`` tuples. Names are those in
|
|
178
|
+
:data:`CDI_ALGORITHMS` (``"dr"`` is an alias of ``"asr"``).
|
|
179
|
+
beta:
|
|
180
|
+
Feedback/relaxation parameter of stages that do not set their own;
|
|
181
|
+
None uses :data:`DEFAULT_BETA`.
|
|
182
|
+
|
|
183
|
+
Returns
|
|
184
|
+
-------
|
|
185
|
+
list of tuple
|
|
186
|
+
One ``(name, iterations, beta)`` per non-empty stage.
|
|
187
|
+
"""
|
|
188
|
+
if isinstance(schedule, str):
|
|
189
|
+
items: list[Sequence[Any]] = []
|
|
190
|
+
for part in schedule.split(","):
|
|
191
|
+
if part.strip():
|
|
192
|
+
items.append([bit.strip() for bit in part.split(":")])
|
|
193
|
+
elif isinstance(schedule, tuple) and schedule and isinstance(schedule[0], str):
|
|
194
|
+
items = [schedule]
|
|
195
|
+
else:
|
|
196
|
+
items = list(schedule)
|
|
197
|
+
stages: list[tuple[str, int, float]] = []
|
|
198
|
+
for item in items:
|
|
199
|
+
if isinstance(item, str) or len(item) not in (2, 3):
|
|
200
|
+
raise ValueError(
|
|
201
|
+
f"schedule stage {item!r} must be (name, iterations) or (name, iterations, beta), "
|
|
202
|
+
"e.g. [('hio', 400), ('er', 100)] or 'hio:400,er:100'"
|
|
203
|
+
)
|
|
204
|
+
name = str(item[0]).strip().lower()
|
|
205
|
+
name = _ALIASES.get(name, name)
|
|
206
|
+
if name not in CDI_ALGORITHMS:
|
|
207
|
+
raise ValueError(f"unknown CDI algorithm {item[0]!r}; choose from {CDI_ALGORITHMS}")
|
|
208
|
+
n = int(item[1])
|
|
209
|
+
if n < 0:
|
|
210
|
+
raise ValueError(f"stage {name!r} has a negative iteration count {n}")
|
|
211
|
+
if len(item) == 3 and item[2] is not None:
|
|
212
|
+
b = float(item[2])
|
|
213
|
+
else:
|
|
214
|
+
b = DEFAULT_BETA[name] if beta is None else float(beta)
|
|
215
|
+
if not b > 0.0:
|
|
216
|
+
raise ValueError(f"stage {name!r} needs beta > 0, got {b}")
|
|
217
|
+
if n:
|
|
218
|
+
stages.append((name, n, b))
|
|
219
|
+
if not stages:
|
|
220
|
+
raise ValueError("the schedule has no iterations; e.g. use schedule='hio:400,er:100'")
|
|
221
|
+
return stages
|
|
222
|
+
|
|
223
|
+
|
|
224
|
+
# ----------------------------------------------------------------------------- result
|
|
225
|
+
@dataclass
|
|
226
|
+
class CDIResult:
|
|
227
|
+
"""Outcome of :func:`cdi`.
|
|
228
|
+
|
|
229
|
+
Array fields are backend arrays (CuPy for GPU solves) until you call
|
|
230
|
+
:meth:`to_numpy`. All images are centred like the inputs.
|
|
231
|
+
|
|
232
|
+
Attributes
|
|
233
|
+
----------
|
|
234
|
+
object:
|
|
235
|
+
``(ny, nx)`` complex object estimate of the best start (the support
|
|
236
|
+
projection of the final data-consistent image).
|
|
237
|
+
support:
|
|
238
|
+
``(ny, nx)`` boolean support of the best start (after shrinkwrap).
|
|
239
|
+
objects:
|
|
240
|
+
``(starts, ny, nx)`` final estimates of every start (for averaging or
|
|
241
|
+
phase-retrieval transfer functions).
|
|
242
|
+
history:
|
|
243
|
+
Fourier-modulus error of the best start at each check; the last entry
|
|
244
|
+
is the final :attr:`modulus_error`.
|
|
245
|
+
support_history:
|
|
246
|
+
Object-domain constraint violation of the best start at the same
|
|
247
|
+
points.
|
|
248
|
+
times:
|
|
249
|
+
Seconds since the start at each :attr:`history` entry.
|
|
250
|
+
n_iter:
|
|
251
|
+
Iterations performed.
|
|
252
|
+
converged:
|
|
253
|
+
Whether the modulus error reached ``tol`` (or the callback stopped).
|
|
254
|
+
message:
|
|
255
|
+
Why the run stopped.
|
|
256
|
+
elapsed:
|
|
257
|
+
Total wall-clock seconds, including setup.
|
|
258
|
+
device:
|
|
259
|
+
``"cpu"`` or ``"gpu"``.
|
|
260
|
+
best_start:
|
|
261
|
+
Index of the start with the lowest final modulus error.
|
|
262
|
+
start_errors:
|
|
263
|
+
Final modulus error of every start (host array).
|
|
264
|
+
modulus_error:
|
|
265
|
+
Final normalized Fourier-modulus error of :attr:`object`,
|
|
266
|
+
``sqrt(sum_meas (|F o| - A)^2 / sum_meas A^2)``.
|
|
267
|
+
support_error:
|
|
268
|
+
Final normalized constraint violation
|
|
269
|
+
``||u - P_S u|| / ||u||`` of the image ``u`` :attr:`object` was
|
|
270
|
+
projected from.
|
|
271
|
+
schedule:
|
|
272
|
+
The parsed schedule ``[(name, iterations, beta), ...]``.
|
|
273
|
+
"""
|
|
274
|
+
|
|
275
|
+
object: Any
|
|
276
|
+
support: Any
|
|
277
|
+
objects: Any
|
|
278
|
+
history: list[float]
|
|
279
|
+
support_history: list[float]
|
|
280
|
+
times: list[float]
|
|
281
|
+
n_iter: int
|
|
282
|
+
converged: bool
|
|
283
|
+
message: str
|
|
284
|
+
elapsed: float
|
|
285
|
+
device: str
|
|
286
|
+
best_start: int
|
|
287
|
+
start_errors: np.ndarray
|
|
288
|
+
modulus_error: float
|
|
289
|
+
support_error: float
|
|
290
|
+
schedule: list[tuple[str, int, float]]
|
|
291
|
+
|
|
292
|
+
def to_numpy(self) -> CDIResult:
|
|
293
|
+
"""Copy with every array field on the host."""
|
|
294
|
+
changes = {
|
|
295
|
+
f.name: to_numpy(getattr(self, f.name))
|
|
296
|
+
for f in fields(self)
|
|
297
|
+
if hasattr(getattr(self, f.name), "shape")
|
|
298
|
+
and not isinstance(getattr(self, f.name), np.ndarray)
|
|
299
|
+
}
|
|
300
|
+
return replace(self, **changes)
|
|
301
|
+
|
|
302
|
+
def __repr__(self) -> str:
|
|
303
|
+
return (
|
|
304
|
+
f"CDIResult(device={self.device!r}, n_iter={self.n_iter}, starts="
|
|
305
|
+
f"{len(self.start_errors)}, best_start={self.best_start}, "
|
|
306
|
+
f"modulus_error={self.modulus_error:.4g}, converged={self.converged})"
|
|
307
|
+
)
|
|
308
|
+
|
|
309
|
+
|
|
310
|
+
# ----------------------------------------------------------------------------- projections
|
|
311
|
+
class _Projector:
|
|
312
|
+
"""Fourier-modulus and object-domain projections on uncentred arrays."""
|
|
313
|
+
|
|
314
|
+
def __init__(
|
|
315
|
+
self,
|
|
316
|
+
be: Backend,
|
|
317
|
+
amplitude: Any,
|
|
318
|
+
measured: Any,
|
|
319
|
+
constraint: str,
|
|
320
|
+
bounds: tuple[float | None, float | None] | None,
|
|
321
|
+
) -> None:
|
|
322
|
+
self.be = be
|
|
323
|
+
self.xp = be.xp
|
|
324
|
+
self.amplitude = amplitude
|
|
325
|
+
self.measured = measured
|
|
326
|
+
self.unmeasured = None if measured is None else ~measured
|
|
327
|
+
self.constraint = constraint
|
|
328
|
+
lo, hi = bounds if bounds is not None else (None, None)
|
|
329
|
+
self.lo = None if lo is None or lo <= 0 else float(lo)
|
|
330
|
+
self.hi = None if hi is None else float(hi)
|
|
331
|
+
peak = float(self.xp.max(amplitude)) if amplitude.size else 0.0
|
|
332
|
+
# Floor for |F x| so a zero transform gives a zero (not NaN) projection.
|
|
333
|
+
self.floor = max(peak * 1e-20, 1e-30)
|
|
334
|
+
weights = amplitude * amplitude
|
|
335
|
+
if measured is not None:
|
|
336
|
+
weights = weights * measured
|
|
337
|
+
self.energy = max(float(self.xp.sum(weights)), 1e-300)
|
|
338
|
+
|
|
339
|
+
def modulus(self, x: Any) -> Any:
|
|
340
|
+
"""``P_M``: impose the measured modulus on measured pixels."""
|
|
341
|
+
xp = self.xp
|
|
342
|
+
spec = self.be.fft2(x)
|
|
343
|
+
factor = xp.abs(spec)
|
|
344
|
+
xp.maximum(factor, self.floor, out=factor)
|
|
345
|
+
xp.divide(self.amplitude, factor, out=factor)
|
|
346
|
+
if self.unmeasured is not None:
|
|
347
|
+
xp.copyto(factor, 1.0, where=self.unmeasured)
|
|
348
|
+
spec *= factor
|
|
349
|
+
return self.be.ifft2(spec)
|
|
350
|
+
|
|
351
|
+
def object(self, u: Any, support: Any) -> Any:
|
|
352
|
+
"""``P_S``: zero outside the support, nearest feasible value inside."""
|
|
353
|
+
xp = self.xp
|
|
354
|
+
if self.constraint == "complex":
|
|
355
|
+
v = u
|
|
356
|
+
elif self.constraint == "real":
|
|
357
|
+
v = u.real
|
|
358
|
+
else:
|
|
359
|
+
v = xp.maximum(u.real, 0.0)
|
|
360
|
+
if self.lo is not None or self.hi is not None:
|
|
361
|
+
if self.constraint == "positive":
|
|
362
|
+
v = xp.clip(v, self.lo, self.hi)
|
|
363
|
+
else:
|
|
364
|
+
mag = xp.abs(v)
|
|
365
|
+
clipped = xp.clip(mag, self.lo, self.hi)
|
|
366
|
+
if self.constraint == "real":
|
|
367
|
+
v = xp.where(v < 0, -clipped, clipped)
|
|
368
|
+
else:
|
|
369
|
+
v = v * (clipped / xp.maximum(mag, self.floor))
|
|
370
|
+
return xp.where(support, v, 0.0)
|
|
371
|
+
|
|
372
|
+
def feasible(self, u: Any, support: Any) -> Any:
|
|
373
|
+
"""Pixels where ``u`` already satisfies the object constraints (for HIO)."""
|
|
374
|
+
ok = support
|
|
375
|
+
if self.constraint == "positive":
|
|
376
|
+
ok = ok & (u.real >= 0.0)
|
|
377
|
+
if self.lo is not None or self.hi is not None:
|
|
378
|
+
mag = self.xp.abs(u) if self.constraint == "complex" else self.xp.abs(u.real)
|
|
379
|
+
if self.lo is not None:
|
|
380
|
+
ok = ok & (mag >= self.lo)
|
|
381
|
+
if self.hi is not None:
|
|
382
|
+
ok = ok & (mag <= self.hi)
|
|
383
|
+
return ok
|
|
384
|
+
|
|
385
|
+
def errors(self, u: Any, support: Any) -> tuple[Any, Any, Any]:
|
|
386
|
+
"""Estimate ``P_S u`` with its per-start modulus error and violation (device arrays)."""
|
|
387
|
+
xp = self.xp
|
|
388
|
+
est = self.object(u, support)
|
|
389
|
+
diff = xp.abs(self.be.fft2(est)) - self.amplitude
|
|
390
|
+
if self.measured is not None:
|
|
391
|
+
diff = diff * self.measured
|
|
392
|
+
mod = xp.sqrt(xp.sum(diff * diff, axis=(-2, -1)) / self.energy)
|
|
393
|
+
resid = u - est
|
|
394
|
+
num = xp.sum(resid.real**2 + resid.imag**2, axis=(-2, -1))
|
|
395
|
+
den = xp.sum(u.real**2 + u.imag**2, axis=(-2, -1))
|
|
396
|
+
viol = xp.sqrt(num / xp.maximum(den, 1e-300))
|
|
397
|
+
return est, mod, viol
|
|
398
|
+
|
|
399
|
+
|
|
400
|
+
def _frequency_grid(be: Backend, shape: tuple[int, int]) -> tuple[Any, Any]:
|
|
401
|
+
"""Uncentred integer frequency indices ``(ky[:, None], kx[None, :])`` on the backend."""
|
|
402
|
+
ny, nx = shape
|
|
403
|
+
ky = np.fft.fftfreq(ny) * ny
|
|
404
|
+
kx = np.fft.fftfreq(nx) * nx
|
|
405
|
+
return be.asarray(ky[:, None], dtype="real"), be.asarray(kx[None, :], dtype="real")
|
|
406
|
+
|
|
407
|
+
|
|
408
|
+
def _gaussian_blur(be: Backend, image: Any, sigma: float, k2: Any) -> Any:
|
|
409
|
+
"""Circular Gaussian blur (std ``sigma`` pixels) of real images over the last two axes."""
|
|
410
|
+
if sigma <= 0:
|
|
411
|
+
return image
|
|
412
|
+
# k2 holds (ky/ny)^2 + (kx/nx)^2 in cycles per pixel squared.
|
|
413
|
+
transfer = be.xp.exp((-2.0 * math.pi**2 * sigma**2) * k2)
|
|
414
|
+
return be.ifft2(be.fft2(image) * transfer).real
|
|
415
|
+
|
|
416
|
+
|
|
417
|
+
def _shrinkwrap(
|
|
418
|
+
be: Backend, image: Any, sigma: float, threshold: float, k2: Any, bounds_mask: Any
|
|
419
|
+
) -> Any:
|
|
420
|
+
xp = be.xp
|
|
421
|
+
blurred = _gaussian_blur(be, xp.abs(image), sigma, k2)
|
|
422
|
+
peak = xp.max(blurred, axis=(-2, -1), keepdims=True)
|
|
423
|
+
support = blurred > threshold * peak
|
|
424
|
+
if bounds_mask is not None:
|
|
425
|
+
support &= bounds_mask
|
|
426
|
+
return support
|
|
427
|
+
|
|
428
|
+
|
|
429
|
+
# ----------------------------------------------------------------------------- the solver
|
|
430
|
+
def cdi(
|
|
431
|
+
magnitudes: Any,
|
|
432
|
+
support: Any,
|
|
433
|
+
*,
|
|
434
|
+
schedule: ScheduleLike = "hio:500,er:100",
|
|
435
|
+
starts: int = 1,
|
|
436
|
+
constraint: str = "complex",
|
|
437
|
+
bounds: tuple[float | None, float | None] | None = None,
|
|
438
|
+
shrinkwrap: ShrinkwrapLike = None,
|
|
439
|
+
measured: Any = None,
|
|
440
|
+
intensity: bool = False,
|
|
441
|
+
beta: float | None = None,
|
|
442
|
+
initial: Any = None,
|
|
443
|
+
seed: Any = None,
|
|
444
|
+
device: BackendLike = "cpu",
|
|
445
|
+
precision: str | None = None,
|
|
446
|
+
check_every: int = 10,
|
|
447
|
+
tol: float = 1e-6,
|
|
448
|
+
oss_stages: int = 10,
|
|
449
|
+
callback: Callable[[int, Any, float], bool | None] | None = None,
|
|
450
|
+
) -> CDIResult:
|
|
451
|
+
"""Reconstruct an object from its far-field modulus by iterative projections.
|
|
452
|
+
|
|
453
|
+
Parameters
|
|
454
|
+
----------
|
|
455
|
+
magnitudes:
|
|
456
|
+
``(ny, nx)`` measured far-field modulus ``|F o|``, centred (DC at
|
|
457
|
+
``(ny // 2, nx // 2)``, i.e. ``fftshift``-ed), any linear units.
|
|
458
|
+
Non-finite values are treated as unmeasured.
|
|
459
|
+
support:
|
|
460
|
+
``(ny, nx)`` boolean object support, centred like the object. Use
|
|
461
|
+
:func:`autocorrelation_support` with ``shrinkwrap`` when it is
|
|
462
|
+
unknown.
|
|
463
|
+
schedule:
|
|
464
|
+
Algorithms to run in order, e.g. ``"hio:400,er:100"`` or
|
|
465
|
+
``[("raar", 300, 0.8), ("er", 50)]``; see :func:`parse_schedule`.
|
|
466
|
+
starts:
|
|
467
|
+
Number of independent random starts run simultaneously as a batch.
|
|
468
|
+
The start with the lowest final modulus error is returned.
|
|
469
|
+
constraint:
|
|
470
|
+
Object-domain constraint inside the support: ``"complex"`` (support
|
|
471
|
+
only), ``"real"`` or ``"positive"`` (real and non-negative).
|
|
472
|
+
bounds:
|
|
473
|
+
Optional ``(min, max)`` bounds on the object modulus inside the
|
|
474
|
+
support (either may be None), in object units.
|
|
475
|
+
shrinkwrap:
|
|
476
|
+
Support refinement: None/False (fixed support), True (defaults), a
|
|
477
|
+
dict of :class:`ShrinkwrapConfig` fields, or a config. Refined
|
|
478
|
+
supports stay inside the initial ``support`` (unless it is all True),
|
|
479
|
+
so pass a generous one, e.g. :func:`autocorrelation_support`.
|
|
480
|
+
measured:
|
|
481
|
+
Optional ``(ny, nx)`` boolean mask, centred, False where the modulus
|
|
482
|
+
is unknown (beamstop, detector gaps); the modulus is left free there.
|
|
483
|
+
intensity:
|
|
484
|
+
If True, ``magnitudes`` holds intensities ``|F o|^2`` (negative values
|
|
485
|
+
are clipped to zero before the square root).
|
|
486
|
+
beta:
|
|
487
|
+
Feedback (HIO, HPR, OSS) / relaxation (DM, RAAR, RRR) parameter of
|
|
488
|
+
stages that do not set their own. None uses the per-algorithm
|
|
489
|
+
:data:`DEFAULT_BETA` (0.9; RAAR 0.98; RRR 0.5). ER and ASR ignore it.
|
|
490
|
+
initial:
|
|
491
|
+
Optional starting object, ``(ny, nx)`` (shared by all starts) or
|
|
492
|
+
``(starts, ny, nx)``, centred. Default: random Fourier phases with the
|
|
493
|
+
measured modulus, restricted to the support.
|
|
494
|
+
seed:
|
|
495
|
+
Seed or :class:`numpy.random.Generator` for the random starts (drawn
|
|
496
|
+
on the host, so CPU and GPU runs start identically).
|
|
497
|
+
device, precision:
|
|
498
|
+
Backend selection, see :func:`solvephase.get_backend`.
|
|
499
|
+
check_every:
|
|
500
|
+
Iterations between error evaluations (each one synchronizes the GPU).
|
|
501
|
+
tol:
|
|
502
|
+
Stop when the best start's modulus error is at or below ``tol``.
|
|
503
|
+
oss_stages:
|
|
504
|
+
Number of filter widths an OSS stage steps through (its iterations
|
|
505
|
+
are split evenly); at the end of each the best iterate is kept.
|
|
506
|
+
callback:
|
|
507
|
+
``callback(iteration, object, error)`` at every check with the
|
|
508
|
+
centred estimate of the currently best start; return True to stop.
|
|
509
|
+
|
|
510
|
+
Returns
|
|
511
|
+
-------
|
|
512
|
+
CDIResult
|
|
513
|
+
The best start's object and support, error histories and per-start
|
|
514
|
+
final errors.
|
|
515
|
+
|
|
516
|
+
Notes
|
|
517
|
+
-----
|
|
518
|
+
Each iteration of ER, HIO, RAAR, RRR, ASR and HPR costs one batched FFT
|
|
519
|
+
pair; DM and OSS cost two. An error check adds one forward FFT.
|
|
520
|
+
"""
|
|
521
|
+
t0 = time.perf_counter()
|
|
522
|
+
be = get_backend(device, precision)
|
|
523
|
+
xp = be.xp
|
|
524
|
+
stages = parse_schedule(schedule, beta)
|
|
525
|
+
constraint = str(constraint).lower()
|
|
526
|
+
if constraint not in _CONSTRAINTS:
|
|
527
|
+
raise ValueError(f"constraint must be one of {_CONSTRAINTS}, got {constraint!r}")
|
|
528
|
+
if int(starts) < 1:
|
|
529
|
+
raise ValueError(f"starts must be >= 1, got {starts}")
|
|
530
|
+
if int(check_every) < 1:
|
|
531
|
+
raise ValueError(f"check_every must be >= 1, got {check_every}")
|
|
532
|
+
if int(oss_stages) < 1:
|
|
533
|
+
raise ValueError(f"oss_stages must be >= 1, got {oss_stages}")
|
|
534
|
+
starts, check_every, oss_stages = int(starts), int(check_every), int(oss_stages)
|
|
535
|
+
sw = _shrinkwrap_config(shrinkwrap)
|
|
536
|
+
|
|
537
|
+
data = np.asarray(to_numpy(magnitudes), dtype=np.float64)
|
|
538
|
+
if data.ndim != 2:
|
|
539
|
+
raise ValueError(f"magnitudes must be a 2-D (ny, nx) array, got shape {data.shape}")
|
|
540
|
+
shape = (int(data.shape[0]), int(data.shape[1]))
|
|
541
|
+
finite = np.isfinite(data)
|
|
542
|
+
data = np.where(finite, data, 0.0)
|
|
543
|
+
if intensity:
|
|
544
|
+
data = np.sqrt(np.maximum(data, 0.0))
|
|
545
|
+
elif np.any(data < 0):
|
|
546
|
+
raise ValueError("magnitudes must be non-negative; pass intensity=True for intensities")
|
|
547
|
+
meas_host = finite.copy()
|
|
548
|
+
if measured is not None:
|
|
549
|
+
m = np.asarray(to_numpy(measured), dtype=bool)
|
|
550
|
+
if m.shape != shape:
|
|
551
|
+
raise ValueError(f"measured has shape {m.shape}, expected {shape}")
|
|
552
|
+
meas_host &= m
|
|
553
|
+
sup_host = np.asarray(to_numpy(support), dtype=bool)
|
|
554
|
+
if sup_host.shape != shape:
|
|
555
|
+
raise ValueError(f"support has shape {sup_host.shape}, expected {shape}")
|
|
556
|
+
if not sup_host.any():
|
|
557
|
+
raise ValueError("the support is empty; pass a boolean mask that covers the object")
|
|
558
|
+
|
|
559
|
+
# One shift into uncentred (FFT) order; everything below is uncentred.
|
|
560
|
+
amp = be.asarray(np.fft.ifftshift(data), dtype="real")
|
|
561
|
+
meas = None if meas_host.all() else be.asarray(np.fft.ifftshift(meas_host))
|
|
562
|
+
if meas is not None:
|
|
563
|
+
amp = amp * meas
|
|
564
|
+
outer = be.asarray(np.fft.ifftshift(sup_host))
|
|
565
|
+
sup = outer
|
|
566
|
+
proj = _Projector(be, amp, meas, constraint, bounds)
|
|
567
|
+
|
|
568
|
+
rng = be.random(seed)
|
|
569
|
+
if initial is None:
|
|
570
|
+
phases = be.asarray(rng.uniform(0.0, 2.0 * math.pi, size=(starts, *shape)), dtype="real")
|
|
571
|
+
x = be.ifft2(amp * xp.exp(1j * phases).astype(be.complex_dtype, copy=False))
|
|
572
|
+
x = xp.where(sup, x, 0.0)
|
|
573
|
+
else:
|
|
574
|
+
init = np.asarray(to_numpy(initial))
|
|
575
|
+
if init.shape == shape:
|
|
576
|
+
init = np.broadcast_to(init, (starts, *shape))
|
|
577
|
+
if init.shape != (starts, *shape):
|
|
578
|
+
raise ValueError(f"initial must have shape {shape} or {(starts, *shape)}")
|
|
579
|
+
x = be.asarray(np.fft.ifftshift(init, axes=(-2, -1)), dtype="complex").copy()
|
|
580
|
+
x = x.astype(be.complex_dtype, copy=False)
|
|
581
|
+
|
|
582
|
+
ky, kx = _frequency_grid(be, shape)
|
|
583
|
+
k2 = (ky / shape[0]) ** 2 + (kx / shape[1]) ** 2
|
|
584
|
+
sw_sigma = sw.sigma if sw is not None else 0.0
|
|
585
|
+
sw_bound = None if (sw is None or sup_host.all()) else outer
|
|
586
|
+
|
|
587
|
+
history: list[float] = []
|
|
588
|
+
viol_history: list[float] = []
|
|
589
|
+
times: list[float] = []
|
|
590
|
+
message, converged = "iteration limit reached", False
|
|
591
|
+
it = 0
|
|
592
|
+
stop = False
|
|
593
|
+
|
|
594
|
+
def report(u: Any) -> tuple[np.ndarray, np.ndarray, Any]:
|
|
595
|
+
est, mod, viol = proj.errors(u, sup)
|
|
596
|
+
mod_h = np.asarray(be.to_numpy(mod), dtype=np.float64)
|
|
597
|
+
viol_h = np.asarray(be.to_numpy(viol), dtype=np.float64)
|
|
598
|
+
best = int(np.argmin(mod_h))
|
|
599
|
+
history.append(float(mod_h[best]))
|
|
600
|
+
viol_history.append(float(viol_h[best]))
|
|
601
|
+
times.append(time.perf_counter() - t0)
|
|
602
|
+
return mod_h, viol_h, est
|
|
603
|
+
|
|
604
|
+
for name, n_stage, b in stages:
|
|
605
|
+
if stop:
|
|
606
|
+
break
|
|
607
|
+
# OSS: the filter width alpha (frequency pixels) steps linearly from 2N to N/5, as
|
|
608
|
+
# in the reference implementation of Rodriguez et al. (2013).
|
|
609
|
+
segment_ends: set[int] = set()
|
|
610
|
+
if name == "oss":
|
|
611
|
+
n_big = float(max(shape))
|
|
612
|
+
alphas = np.linspace(2.0 * n_big, 0.2 * n_big, oss_stages)
|
|
613
|
+
edges = np.linspace(0, n_stage, oss_stages + 1).round().astype(int)
|
|
614
|
+
segment_ends = {int(e) for e in edges[1:]}
|
|
615
|
+
kk = ky**2 + kx**2
|
|
616
|
+
best_x, top_x = x.copy(), x.copy()
|
|
617
|
+
best_err: np.ndarray = np.full(starts, np.inf)
|
|
618
|
+
top_err: np.ndarray = np.full(starts, np.inf)
|
|
619
|
+
segment = 0
|
|
620
|
+
window = xp.exp(-0.5 * kk / float(alphas[0]) ** 2)
|
|
621
|
+
for local in range(1, n_stage + 1):
|
|
622
|
+
it += 1
|
|
623
|
+
check = it % check_every == 0 or local in segment_ends
|
|
624
|
+
x_old = x
|
|
625
|
+
y = proj.modulus(x)
|
|
626
|
+
u = y
|
|
627
|
+
if name == "er":
|
|
628
|
+
x = proj.object(y, sup)
|
|
629
|
+
elif name in ("hio", "oss"):
|
|
630
|
+
x = xp.where(proj.feasible(y, sup), proj.object(y, sup), x - b * y)
|
|
631
|
+
if name == "oss":
|
|
632
|
+
smooth = be.ifft2(be.fft2(x) * window)
|
|
633
|
+
x = xp.where(sup, x, smooth)
|
|
634
|
+
elif name == "hpr":
|
|
635
|
+
x = x - b * y + proj.object((1.0 + b) * y - x, sup)
|
|
636
|
+
elif name == "asr":
|
|
637
|
+
x = x - y + proj.object(2.0 * y - x, sup)
|
|
638
|
+
elif name == "rrr":
|
|
639
|
+
x = x + b * (proj.object(2.0 * y - x, sup) - y)
|
|
640
|
+
elif name == "raar":
|
|
641
|
+
x = b * x + (1.0 - 2.0 * b) * y + b * proj.object(2.0 * y - x, sup)
|
|
642
|
+
else: # dm, Elser 2003 with gamma_S = -1/beta, gamma_M = 1/beta
|
|
643
|
+
g_s, g_m = -1.0 / b, 1.0 / b
|
|
644
|
+
u = (1.0 + g_m) * y - g_m * x
|
|
645
|
+
f_s = (1.0 + g_s) * proj.object(x, sup) - g_s * x
|
|
646
|
+
x = x + b * (proj.object(u, sup) - proj.modulus(f_s))
|
|
647
|
+
x = x.astype(be.complex_dtype, copy=False)
|
|
648
|
+
|
|
649
|
+
if check:
|
|
650
|
+
mod_h, _, est = report(u)
|
|
651
|
+
if name == "oss":
|
|
652
|
+
# Keep each start's best iterate of the segment; restart the next
|
|
653
|
+
# segment from it, and end the stage on the best of all segments.
|
|
654
|
+
better = mod_h < best_err
|
|
655
|
+
if better.any():
|
|
656
|
+
best_x = xp.where(be.asarray(better)[:, None, None], x_old, best_x)
|
|
657
|
+
best_err = np.where(better, mod_h, best_err)
|
|
658
|
+
if local in segment_ends:
|
|
659
|
+
top = best_err < top_err
|
|
660
|
+
top_x = xp.where(be.asarray(top)[:, None, None], best_x, top_x)
|
|
661
|
+
top_err = np.where(top, best_err, top_err)
|
|
662
|
+
x = top_x.copy() if local == n_stage else best_x.copy()
|
|
663
|
+
best_err = np.full(starts, np.inf)
|
|
664
|
+
segment += 1
|
|
665
|
+
if segment < len(alphas):
|
|
666
|
+
window = xp.exp(-0.5 * kk / float(alphas[segment]) ** 2)
|
|
667
|
+
if callback is not None:
|
|
668
|
+
best = int(np.argmin(mod_h))
|
|
669
|
+
centred = xp.fft.fftshift(est[best], axes=(-2, -1))
|
|
670
|
+
if callback(it, centred, float(mod_h[best])):
|
|
671
|
+
message, converged, stop = "stopped by callback", True, True
|
|
672
|
+
if not stop and float(mod_h.min()) <= tol:
|
|
673
|
+
message, converged, stop = "modulus error reached tol", True, True
|
|
674
|
+
if (
|
|
675
|
+
sw is not None
|
|
676
|
+
and it % sw.every == 0
|
|
677
|
+
and it >= sw.start
|
|
678
|
+
and (sw.stop is None or it <= sw.stop)
|
|
679
|
+
):
|
|
680
|
+
sup = _shrinkwrap(be, y, sw_sigma, sw.threshold, k2, sw_bound)
|
|
681
|
+
sw_sigma = max(sw.sigma_min, sw_sigma * sw.decay)
|
|
682
|
+
if stop:
|
|
683
|
+
break
|
|
684
|
+
|
|
685
|
+
# Final estimate: the support projection of the data-consistent image of x
|
|
686
|
+
# (for DM, of f_M(x)); it is also the last history entry.
|
|
687
|
+
y = proj.modulus(x)
|
|
688
|
+
if name == "dm":
|
|
689
|
+
y = (1.0 + 1.0 / b) * y - (1.0 / b) * x
|
|
690
|
+
mod_h, viol_h, est = report(y)
|
|
691
|
+
best = int(np.argmin(mod_h))
|
|
692
|
+
objects = xp.fft.fftshift(est.astype(be.complex_dtype, copy=False), axes=(-2, -1))
|
|
693
|
+
sup_b = sup if sup.ndim == 2 else sup[best]
|
|
694
|
+
be.synchronize()
|
|
695
|
+
return CDIResult(
|
|
696
|
+
object=objects[best],
|
|
697
|
+
support=xp.fft.fftshift(sup_b, axes=(-2, -1)),
|
|
698
|
+
objects=objects,
|
|
699
|
+
history=history,
|
|
700
|
+
support_history=viol_history,
|
|
701
|
+
times=times,
|
|
702
|
+
n_iter=it,
|
|
703
|
+
converged=converged,
|
|
704
|
+
message=message,
|
|
705
|
+
elapsed=time.perf_counter() - t0,
|
|
706
|
+
device=be.device,
|
|
707
|
+
best_start=best,
|
|
708
|
+
start_errors=mod_h,
|
|
709
|
+
modulus_error=float(mod_h[best]),
|
|
710
|
+
support_error=float(viol_h[best]),
|
|
711
|
+
schedule=stages,
|
|
712
|
+
)
|
|
713
|
+
|
|
714
|
+
|
|
715
|
+
# ----------------------------------------------------------------------------- helpers
|
|
716
|
+
def autocorrelation_support(
|
|
717
|
+
intensity: Any,
|
|
718
|
+
threshold: float = 0.04,
|
|
719
|
+
*,
|
|
720
|
+
measured: Any = None,
|
|
721
|
+
device: BackendLike = None,
|
|
722
|
+
) -> Any:
|
|
723
|
+
"""Loose initial support from the thresholded autocorrelation.
|
|
724
|
+
|
|
725
|
+
The inverse Fourier transform of the far-field intensity is the object's
|
|
726
|
+
autocorrelation, whose support is twice the object's extent (the
|
|
727
|
+
difference set ``S - S``); thresholded, it is the usual starting support
|
|
728
|
+
for shrinkwrap.
|
|
729
|
+
|
|
730
|
+
Parameters
|
|
731
|
+
----------
|
|
732
|
+
intensity:
|
|
733
|
+
``(ny, nx)`` centred far-field intensity ``|F o|^2`` (square
|
|
734
|
+
magnitudes first). Non-finite values count as unmeasured.
|
|
735
|
+
threshold:
|
|
736
|
+
Fraction of the autocorrelation maximum kept.
|
|
737
|
+
measured:
|
|
738
|
+
Optional centred boolean mask, False where the intensity is unknown
|
|
739
|
+
(set to zero before transforming).
|
|
740
|
+
device:
|
|
741
|
+
Backend; by default the backend of ``intensity``.
|
|
742
|
+
|
|
743
|
+
Returns
|
|
744
|
+
-------
|
|
745
|
+
array
|
|
746
|
+
``(ny, nx)`` centred boolean support on the backend.
|
|
747
|
+
"""
|
|
748
|
+
if not 0.0 < threshold < 1.0:
|
|
749
|
+
raise ValueError(f"threshold must be in (0, 1), got {threshold}")
|
|
750
|
+
be = backend_of(intensity) if device is None else get_backend(device)
|
|
751
|
+
xp = be.xp
|
|
752
|
+
data = be.asarray(intensity, dtype="real")
|
|
753
|
+
if data.ndim != 2:
|
|
754
|
+
raise ValueError(f"intensity must be 2-D, got shape {tuple(data.shape)}")
|
|
755
|
+
data = xp.where(xp.isfinite(data), data, 0.0)
|
|
756
|
+
if measured is not None:
|
|
757
|
+
data = data * be.asarray(measured, dtype=bool)
|
|
758
|
+
auto = xp.abs(xp.fft.fftshift(be.ifft2(xp.fft.ifftshift(data))))
|
|
759
|
+
return auto > threshold * xp.max(auto)
|
|
760
|
+
|
|
761
|
+
|
|
762
|
+
def _upsampled_correlation(
|
|
763
|
+
xp: Any, product: Any, centre: tuple[float, float], upsample: int, half: int
|
|
764
|
+
) -> Any:
|
|
765
|
+
"""Cross-correlation sampled on ``centre + m / upsample``, ``|m| <= half``.
|
|
766
|
+
|
|
767
|
+
``product`` is ``F(ref) * conj(F(est))`` (uncentred); the correlation at
|
|
768
|
+
shift ``s`` is ``sum_k product(k) exp(+2 pi i k.s / n)``
|
|
769
|
+
(Guizar-Sicairos et al. 2008, matrix-multiply DFT).
|
|
770
|
+
"""
|
|
771
|
+
ny, nx = product.shape
|
|
772
|
+
offsets = xp.arange(-half, half + 1) / float(upsample)
|
|
773
|
+
ky = xp.asarray(np.fft.fftfreq(ny) * ny)
|
|
774
|
+
kx = xp.asarray(np.fft.fftfreq(nx) * nx)
|
|
775
|
+
ey = xp.exp((2j * math.pi / ny) * xp.outer(centre[0] + offsets, ky))
|
|
776
|
+
ex = xp.exp((2j * math.pi / nx) * xp.outer(kx, centre[1] + offsets))
|
|
777
|
+
return ey @ product @ ex
|
|
778
|
+
|
|
779
|
+
|
|
780
|
+
def align_object(
|
|
781
|
+
estimate: Any,
|
|
782
|
+
reference: Any,
|
|
783
|
+
*,
|
|
784
|
+
subpixel: bool = True,
|
|
785
|
+
upsample: int = 100,
|
|
786
|
+
twin: bool = True,
|
|
787
|
+
scale: bool = True,
|
|
788
|
+
) -> tuple[Any, float]:
|
|
789
|
+
"""Register a CDI reconstruction to a reference, removing the trivial ambiguities.
|
|
790
|
+
|
|
791
|
+
A far-field modulus does not change under a global phase factor, a
|
|
792
|
+
translation, or the twin image ``conj(o(-r))``. This finds the
|
|
793
|
+
translation (integer cross-correlation peak refined to ``1/upsample``
|
|
794
|
+
pixel by an upsampled matrix DFT, Guizar-Sicairos et al. 2008), the twin
|
|
795
|
+
choice and the complex factor ``c`` that minimize
|
|
796
|
+
``||c * shift(est) - ref|| / ||ref||``.
|
|
797
|
+
|
|
798
|
+
Parameters
|
|
799
|
+
----------
|
|
800
|
+
estimate, reference:
|
|
801
|
+
``(ny, nx)`` complex or real images of the same shape.
|
|
802
|
+
subpixel:
|
|
803
|
+
Refine the translation below one pixel (applied as a Fourier phase
|
|
804
|
+
ramp, i.e. a circular shift).
|
|
805
|
+
upsample:
|
|
806
|
+
Sub-pixel refinement factor.
|
|
807
|
+
twin:
|
|
808
|
+
Also try the twin ``conj(est[::-1, ::-1])`` and keep the better one.
|
|
809
|
+
scale:
|
|
810
|
+
If True ``c`` is any complex number (phase and scale); if False only
|
|
811
|
+
a global phase ``|c| = 1`` is removed.
|
|
812
|
+
|
|
813
|
+
Returns
|
|
814
|
+
-------
|
|
815
|
+
aligned:
|
|
816
|
+
``c * shift(est)`` (or its twin), on the backend of ``estimate``.
|
|
817
|
+
error:
|
|
818
|
+
``||aligned - ref|| / ||ref||``.
|
|
819
|
+
"""
|
|
820
|
+
be = backend_of(estimate)
|
|
821
|
+
if be.precision == "single":
|
|
822
|
+
be = get_backend(be.device, "double")
|
|
823
|
+
xp = be.xp
|
|
824
|
+
est = be.asarray(estimate, dtype="complex")
|
|
825
|
+
ref = be.asarray(reference, dtype="complex")
|
|
826
|
+
if est.ndim != 2 or est.shape != ref.shape:
|
|
827
|
+
raise ValueError(
|
|
828
|
+
f"estimate and reference must be 2-D with the same shape, got "
|
|
829
|
+
f"{tuple(est.shape)} and {tuple(ref.shape)}"
|
|
830
|
+
)
|
|
831
|
+
ny, nx = est.shape
|
|
832
|
+
f_ref = be.fft2(ref)
|
|
833
|
+
ref_norm = math.sqrt(be.dot(ref, ref))
|
|
834
|
+
if ref_norm == 0:
|
|
835
|
+
raise ValueError("the reference is zero")
|
|
836
|
+
ky, kx = _frequency_grid(be, (ny, nx))
|
|
837
|
+
candidates = [est]
|
|
838
|
+
if twin:
|
|
839
|
+
candidates.append(xp.conj(est[::-1, ::-1]))
|
|
840
|
+
best: tuple[Any, float] | None = None
|
|
841
|
+
for cand in candidates:
|
|
842
|
+
f_c = be.fft2(cand)
|
|
843
|
+
product = f_ref * xp.conj(f_c)
|
|
844
|
+
corr = xp.abs(be.ifft2(product))
|
|
845
|
+
iy, ix = (int(v) for v in np.unravel_index(int(xp.argmax(corr)), corr.shape))
|
|
846
|
+
sy = float(iy - ny if iy > ny // 2 else iy)
|
|
847
|
+
sx = float(ix - nx if ix > nx // 2 else ix)
|
|
848
|
+
if subpixel and upsample > 1:
|
|
849
|
+
half = math.ceil(0.75 * upsample)
|
|
850
|
+
local = xp.abs(_upsampled_correlation(xp, product, (sy, sx), int(upsample), half))
|
|
851
|
+
my, mx = np.unravel_index(int(xp.argmax(local)), local.shape)
|
|
852
|
+
sy += (int(my) - half) / float(upsample)
|
|
853
|
+
sx += (int(mx) - half) / float(upsample)
|
|
854
|
+
ramp = xp.exp((-2j * math.pi) * (ky * (sy / ny) + kx * (sx / nx)))
|
|
855
|
+
shifted = be.ifft2(f_c * ramp)
|
|
856
|
+
if sy == round(sy) and sx == round(sx):
|
|
857
|
+
shifted = xp.roll(cand, (round(sy), round(sx)), axis=(0, 1))
|
|
858
|
+
energy = be.dot(shifted, shifted)
|
|
859
|
+
if energy == 0:
|
|
860
|
+
c: complex = 0.0
|
|
861
|
+
else:
|
|
862
|
+
c = complex(xp.sum(xp.conj(shifted) * ref)) / energy
|
|
863
|
+
if not scale:
|
|
864
|
+
c = c / abs(c) if c != 0 else 1.0
|
|
865
|
+
aligned = c * shifted
|
|
866
|
+
resid = aligned - ref
|
|
867
|
+
err = math.sqrt(be.dot(resid, resid)) / ref_norm
|
|
868
|
+
if best is None or err < best[1]:
|
|
869
|
+
best = (aligned, err)
|
|
870
|
+
assert best is not None
|
|
871
|
+
return best
|
|
872
|
+
|
|
873
|
+
|
|
874
|
+
@dataclass
|
|
875
|
+
class SimulatedCDI:
|
|
876
|
+
"""Synthetic CDI data from :func:`simulate_cdi` (centred backend arrays).
|
|
877
|
+
|
|
878
|
+
Attributes
|
|
879
|
+
----------
|
|
880
|
+
magnitudes:
|
|
881
|
+
``(ny, nx)`` far-field modulus: ``|F o|`` without noise, or the square
|
|
882
|
+
root of the Poisson counts.
|
|
883
|
+
intensity:
|
|
884
|
+
``(ny, nx)`` far-field intensity: ``|F o|^2`` or the Poisson counts.
|
|
885
|
+
object:
|
|
886
|
+
``(ny, nx)`` zero-padded true object, scaled so that ``|F o|^2`` is the
|
|
887
|
+
expected intensity.
|
|
888
|
+
support:
|
|
889
|
+
``(ny, nx)`` tight support ``|o| > 0`` of the padded object.
|
|
890
|
+
"""
|
|
891
|
+
|
|
892
|
+
magnitudes: Any
|
|
893
|
+
intensity: Any
|
|
894
|
+
object: Any
|
|
895
|
+
support: Any
|
|
896
|
+
|
|
897
|
+
|
|
898
|
+
def simulate_cdi(
|
|
899
|
+
obj: Any,
|
|
900
|
+
oversampling: float = 2.0,
|
|
901
|
+
*,
|
|
902
|
+
shape: tuple[int, int] | None = None,
|
|
903
|
+
photons: float | None = None,
|
|
904
|
+
seed: Any = None,
|
|
905
|
+
device: BackendLike = "cpu",
|
|
906
|
+
precision: str | None = None,
|
|
907
|
+
) -> SimulatedCDI:
|
|
908
|
+
"""Zero-pad an object and compute its centred far-field diffraction data.
|
|
909
|
+
|
|
910
|
+
Parameters
|
|
911
|
+
----------
|
|
912
|
+
obj:
|
|
913
|
+
``(my, mx)`` complex or real object.
|
|
914
|
+
oversampling:
|
|
915
|
+
Linear oversampling per axis: the window is ``round(oversampling * m)``
|
|
916
|
+
pixels along each axis. Unique recovery needs Miao's oversampling
|
|
917
|
+
ratio (window area / object area) above 2, i.e. ``oversampling``
|
|
918
|
+
above ``sqrt(2)`` for an object filling its box.
|
|
919
|
+
shape:
|
|
920
|
+
Explicit ``(ny, nx)`` window instead of ``oversampling``.
|
|
921
|
+
photons:
|
|
922
|
+
If given, the expected total photon count: the intensity is scaled
|
|
923
|
+
to it and Poisson noise is drawn (on the host).
|
|
924
|
+
seed:
|
|
925
|
+
Seed or :class:`numpy.random.Generator` for the noise.
|
|
926
|
+
device, precision:
|
|
927
|
+
Backend of the returned arrays.
|
|
928
|
+
|
|
929
|
+
Returns
|
|
930
|
+
-------
|
|
931
|
+
SimulatedCDI
|
|
932
|
+
"""
|
|
933
|
+
be = get_backend(device, precision)
|
|
934
|
+
xp = be.xp
|
|
935
|
+
o = np.asarray(to_numpy(obj))
|
|
936
|
+
if o.ndim != 2:
|
|
937
|
+
raise ValueError(f"obj must be 2-D, got shape {o.shape}")
|
|
938
|
+
my, mx = o.shape
|
|
939
|
+
if shape is None:
|
|
940
|
+
if oversampling < 1:
|
|
941
|
+
raise ValueError(f"oversampling must be >= 1, got {oversampling}")
|
|
942
|
+
shape = (round(oversampling * my), round(oversampling * mx))
|
|
943
|
+
ny, nx = int(shape[0]), int(shape[1])
|
|
944
|
+
if ny < my or nx < mx:
|
|
945
|
+
raise ValueError(f"window {shape} is smaller than the object {o.shape}")
|
|
946
|
+
padded = np.zeros((ny, nx), dtype=np.complex128)
|
|
947
|
+
y0, x0 = (ny - my) // 2, (nx - mx) // 2
|
|
948
|
+
padded[y0 : y0 + my, x0 : x0 + mx] = o
|
|
949
|
+
field = np.fft.fftshift(np.fft.fft2(np.fft.ifftshift(padded), norm="ortho"))
|
|
950
|
+
inten = np.abs(field) ** 2
|
|
951
|
+
if photons is not None:
|
|
952
|
+
if photons <= 0:
|
|
953
|
+
raise ValueError(f"photons must be positive, got {photons}")
|
|
954
|
+
factor = float(photons) / float(inten.sum())
|
|
955
|
+
padded *= math.sqrt(factor)
|
|
956
|
+
expected = inten * factor
|
|
957
|
+
inten = be.random(seed).poisson(expected).astype(np.float64)
|
|
958
|
+
mags = np.sqrt(inten)
|
|
959
|
+
else:
|
|
960
|
+
mags = np.abs(field)
|
|
961
|
+
obj_out = padded if np.iscomplexobj(o) else padded.real
|
|
962
|
+
return SimulatedCDI(
|
|
963
|
+
magnitudes=be.asarray(mags, dtype="real"),
|
|
964
|
+
intensity=be.asarray(inten, dtype="real"),
|
|
965
|
+
object=be.asarray(obj_out),
|
|
966
|
+
support=xp.asarray(np.abs(padded) > 0),
|
|
967
|
+
)
|