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.
@@ -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
+ )