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,693 @@
1
+ """Transport-of-intensity (TIE) phase retrieval from a defocus stack.
2
+
3
+ In the paraxial regime the axial change of intensity of a monochromatic beam
4
+ is fixed by its transverse phase (Teague, JOSA 73, 1434, 1983)::
5
+
6
+ -k dI/dz = div(I grad(phi)), k = 2 pi / lambda.
7
+
8
+ Given intensities ``I(z_j)`` near the plane ``z = 0`` the solver estimates the
9
+ axial derivative ``dI/dz`` and the in-focus intensity ``I0``, then inverts the
10
+ elliptic equation for ``phi`` (radians at ``z = 0``):
11
+
12
+ * **Axial derivative.** ``"central"`` uses the nearest planes on each side of
13
+ ``z = 0`` (and ``z = 0`` itself when measured): a central difference for a
14
+ symmetric pair, the three-point Lagrange derivative otherwise.
15
+ ``"polyfit"`` fits a polynomial of degree ``order`` in ``z`` to every plane
16
+ by least squares, pixel by pixel (Savitzky-Golay style), and takes its
17
+ slope at ``z = 0``; with more planes than coefficients it averages noise,
18
+ with ``order = N - 1`` it is the exact ``N``-point stencil whose truncation
19
+ error falls as a higher power of the plane spacing (Waller et al., Opt.
20
+ Express 18, 12552, 2010). The measured ``z = 0`` plane is the in-focus
21
+ intensity when present; otherwise the interpolant (or fit) value at
22
+ ``z = 0`` is used.
23
+ * **Spectral solvers.** ``"fft"`` assumes a periodic field (real FFTs);
24
+ ``"dct"`` assumes zero normal slope at the window edge and expands the field
25
+ in the cosine series of its mirror extension (type-II DCT; gradients are
26
+ type-II DST series), which avoids the wrap-around coupling of opposite edges
27
+ for non-periodic fields (Zuo, Chen & Asundi, Opt. Express 22, 9220, 2014).
28
+ Both differentiate exactly and use either the uniform-intensity
29
+ approximation ``laplacian(phi) = -k dI/dz / mean(I0)`` (``uniform=True``) or
30
+ Teague's auxiliary function ``psi`` with ``grad(psi) = I grad(phi)``: two
31
+ Poisson solves and a division by the clipped intensity (Gureyev & Nugent,
32
+ JOSA A 13, 1670, 1996; Paganin & Nugent, PRL 80, 2586, 1998).
33
+ * **Masked exact solver.** Teague's construction assumes ``I grad(phi)`` is
34
+ curl free. ``"pcg"`` drops that assumption and solves the five-point
35
+ discretization of ``div(I grad(phi)) = -k dI/dz`` inside a support ``mask``
36
+ (default: where ``I0`` exceeds ``intensity_floor`` of its maximum) by
37
+ conjugate gradients preconditioned with the DCT Poisson solver. Edges that
38
+ leave the mask carry no flux, so it is the right solver for an illuminated
39
+ aperture with dark surroundings, where the edge signal carries the
40
+ boundary slopes (curvature sensing, Roddier, Appl. Opt. 27, 1223, 1988).
41
+ * **Regularization.** The first Poisson solve of ``"fft"``/``"dct"`` uses the
42
+ Tikhonov filter ``lam / (lam**2 + alpha)`` on the eigenvalues ``lam`` of
43
+ ``-laplacian`` with ``alpha = regularization * lam_min**2``: the lowest
44
+ non-zero spatial frequency of the grid is attenuated by exactly
45
+ ``1 / (1 + regularization)``. Raise it to suppress the low-frequency
46
+ "cloud" artefacts that noise produces (Zuo et al., Opt. Lasers Eng. 135,
47
+ 106187, 2020). Teague's second solve is bounded and left unregularized.
48
+
49
+ Piston is not measurable and is removed (mean over the mask, or over each
50
+ connected region of the mask for ``"pcg"``).
51
+ """
52
+
53
+ from __future__ import annotations
54
+
55
+ import functools
56
+ import math
57
+ import time
58
+ from collections.abc import Sequence
59
+ from dataclasses import dataclass, field, fields, replace
60
+ from typing import Any, Literal
61
+
62
+ import numpy as np
63
+
64
+ from ..backend import Backend, BackendLike, _cpu_workers, backend_of, get_backend, to_numpy
65
+ from ..propagation import AngularSpectrumPropagator
66
+
67
+ __all__ = ["TIEResult", "simulate_defocus_stack", "tie"]
68
+
69
+ TIEMethod = Literal["dct", "fft", "pcg"]
70
+ TIEDerivative = Literal["central", "polyfit"]
71
+
72
+
73
+ @dataclass
74
+ class TIEResult:
75
+ """Outcome of a transport-of-intensity retrieval.
76
+
77
+ Array fields are backend arrays (CuPy for GPU solves) until you call
78
+ :meth:`to_numpy`.
79
+
80
+ Attributes
81
+ ----------
82
+ phase:
83
+ ``(ny, nx)`` phase at ``z = 0`` in radians, piston removed, zero
84
+ outside ``mask``.
85
+ opd:
86
+ The same wavefront as optical path difference in metres
87
+ (``phase * wavelength / (2 pi)``).
88
+ intensity:
89
+ ``(ny, nx)`` in-focus intensity ``I0`` used by the solver (measured or
90
+ estimated), in data units.
91
+ didz:
92
+ ``(ny, nx)`` axial intensity derivative estimate, data units per metre.
93
+ mask:
94
+ Boolean support the phase is defined on.
95
+ wavelength:
96
+ Wavelength in metres.
97
+ method:
98
+ Solver: ``"fft"``, ``"dct"`` or ``"pcg"``.
99
+ derivative:
100
+ Axial derivative estimator: ``"central"`` or ``"polyfit"``.
101
+ uniform:
102
+ Whether the uniform-intensity approximation was used.
103
+ history:
104
+ Relative residual norm per convergence check (``"pcg"`` only).
105
+ n_iter:
106
+ Conjugate-gradient iterations (0 for the direct solvers).
107
+ converged:
108
+ Whether the solve finished (always true for direct solvers).
109
+ message:
110
+ Why it stopped.
111
+ elapsed:
112
+ Wall-clock seconds, including the derivative estimate.
113
+ device:
114
+ ``"cpu"`` or ``"gpu"``.
115
+ """
116
+
117
+ phase: Any
118
+ opd: Any
119
+ intensity: Any
120
+ didz: Any
121
+ mask: Any
122
+ wavelength: float
123
+ method: str
124
+ derivative: str
125
+ uniform: bool
126
+ history: list[float] = field(default_factory=list)
127
+ n_iter: int = 0
128
+ converged: bool = True
129
+ message: str = ""
130
+ elapsed: float = 0.0
131
+ device: str = "cpu"
132
+
133
+ def to_numpy(self) -> TIEResult:
134
+ """Copy with every array field on the host."""
135
+ changes: dict[str, Any] = {}
136
+ for f in fields(self):
137
+ value = getattr(self, f.name)
138
+ if hasattr(value, "shape") and not isinstance(value, np.ndarray):
139
+ changes[f.name] = to_numpy(value)
140
+ return replace(self, **changes)
141
+
142
+
143
+ # ----------------------------------------------------------- forward model
144
+ def simulate_defocus_stack(
145
+ field: Any,
146
+ distances: float | Sequence[float],
147
+ pitch: float,
148
+ wavelength: float,
149
+ *,
150
+ paraxial: bool = False,
151
+ device: BackendLike = None,
152
+ precision: str | None = None,
153
+ ) -> Any:
154
+ """Intensities of a complex field propagated to several distances.
155
+
156
+ Uses :class:`~solvephase.propagation.AngularSpectrumPropagator` (periodic
157
+ boundaries: pad the field if light must not wrap around).
158
+
159
+ Parameters
160
+ ----------
161
+ field:
162
+ ``(ny, nx)`` complex field at ``z = 0`` (amplitude, not intensity).
163
+ distances:
164
+ Propagation distances in metres (positive downstream); ``0`` returns
165
+ ``|field|**2``.
166
+ pitch:
167
+ Sample pitch in metres.
168
+ wavelength:
169
+ Wavelength in metres.
170
+ paraxial:
171
+ Use the Fresnel transfer function instead of the exact angular
172
+ spectrum.
173
+ device, precision:
174
+ Backend; defaults to the backend of ``field``.
175
+
176
+ Returns
177
+ -------
178
+ ``(Z, ny, nx)`` intensity stack on the backend (``(ny, nx)`` for a scalar
179
+ distance).
180
+ """
181
+ be = _resolve_backend(field, device, precision)
182
+ u = be.asarray(field, dtype="complex")
183
+ if u.ndim != 2:
184
+ raise ValueError(f"field must be (ny, nx), got shape {tuple(u.shape)}")
185
+ prop = AngularSpectrumPropagator(
186
+ u.shape, pitch, wavelength, distances, paraxial=paraxial, backend=be
187
+ )
188
+ out = prop.forward(u)
189
+ return (out.real**2 + out.imag**2).astype(be.real_dtype, copy=False)
190
+
191
+
192
+ # ------------------------------------------------------- axial derivative
193
+ def _poly_weights(z: np.ndarray, order: int) -> tuple[np.ndarray, np.ndarray]:
194
+ """Least-squares polynomial weights for the value and slope at ``z = 0``."""
195
+ if order < 1:
196
+ raise ValueError(f"order must be at least 1, got {order}")
197
+ if z.size < order + 1:
198
+ raise ValueError(
199
+ f"a degree-{order} fit needs at least {order + 1} planes, got {z.size}; "
200
+ "lower order or add planes"
201
+ )
202
+ scale = float(np.max(np.abs(z)))
203
+ vander = (z[:, None] / scale) ** np.arange(order + 1)[None, :]
204
+ if np.linalg.matrix_rank(vander) < order + 1:
205
+ raise ValueError(f"the plane distances {z.tolist()} cannot constrain a degree-{order} fit")
206
+ pinv = np.linalg.pinv(vander)
207
+ return pinv[0], pinv[1] / scale
208
+
209
+
210
+ def _axial_weights(
211
+ z: np.ndarray, derivative: str, order: int | None
212
+ ) -> tuple[np.ndarray, np.ndarray, int | None]:
213
+ """Plane weights for ``dI/dz`` and ``I0`` at ``z = 0``, and the ``z = 0`` index."""
214
+ n = z.size
215
+ tiny = 1e-9 * float(np.max(np.abs(z)))
216
+ zero = np.flatnonzero(np.abs(z) <= tiny)
217
+ focus = int(zero[0]) if zero.size else None
218
+ w_value = np.zeros(n)
219
+ w_slope = np.zeros(n)
220
+ if derivative == "central":
221
+ if order is not None:
222
+ raise ValueError("order applies to derivative='polyfit' only")
223
+ below = np.flatnonzero(z < -tiny)
224
+ above = np.flatnonzero(z > tiny)
225
+ if below.size == 0 or above.size == 0:
226
+ raise ValueError(
227
+ "derivative='central' needs a plane on each side of z = 0; "
228
+ "use derivative='polyfit' for a one-sided stack"
229
+ )
230
+ idx = [int(below[np.argmax(z[below])]), int(above[np.argmin(z[above])])]
231
+ if focus is not None:
232
+ idx.insert(1, focus)
233
+ sel = np.array(idx)
234
+ value, slope = _poly_weights(z[sel], sel.size - 1)
235
+ w_value[sel] = value
236
+ w_slope[sel] = slope
237
+ elif derivative == "polyfit":
238
+ if order is None:
239
+ order = min(n - 1, 3)
240
+ w_value, w_slope = _poly_weights(z, int(order))
241
+ else:
242
+ raise ValueError(f"derivative must be 'central' or 'polyfit', got {derivative!r}")
243
+ return w_slope, w_value, focus
244
+
245
+
246
+ # ----------------------------------------------------------- Poisson kernels
247
+ def _inverse_filter(be: Backend, lam: np.ndarray, alpha: float) -> Any:
248
+ """Tikhonov-regularized inverse ``lam / (lam**2 + alpha)``, zero mode removed."""
249
+ safe = np.where(lam == 0, 1.0, lam)
250
+ inv = np.where(lam == 0, 0.0, safe / (safe**2 + alpha))
251
+ return be.xp.asarray(inv, dtype=be.real_dtype)
252
+
253
+
254
+ def _lam_min(lams: Sequence[np.ndarray]) -> float:
255
+ """Smallest non-zero eigenvalue over the per-axis eigenvalue sets."""
256
+ return min(float(lam[1]) for lam in lams if lam.size > 1)
257
+
258
+
259
+ @functools.lru_cache(maxsize=4)
260
+ def _fft_plan(
261
+ be: Backend, shape: tuple[int, int], pitch: tuple[float, float], regularization: float
262
+ ) -> tuple[Any, Any, Any, Any]:
263
+ """Wavenumbers ``(qy, qx)`` on the ``rfft2`` grid; regularized and exact inverse ``-lap``."""
264
+ ny, nx = shape
265
+ qy = 2.0 * math.pi * np.fft.fftfreq(ny, d=pitch[0])
266
+ qx = 2.0 * math.pi * np.fft.rfftfreq(nx, d=pitch[1])
267
+ lam = qy[:, None] ** 2 + qx[None, :] ** 2
268
+ alpha = regularization * _lam_min([qy**2, qx**2]) ** 2
269
+ # The Nyquist derivative of a real signal is undefined; drop it.
270
+ if ny % 2 == 0:
271
+ qy[ny // 2] = 0.0
272
+ if nx % 2 == 0:
273
+ qx[-1] = 0.0
274
+ xp, dt = be.xp, be.real_dtype
275
+ qy_d, qx_d = xp.asarray(qy[:, None], dtype=dt), xp.asarray(qx[None, :], dtype=dt)
276
+ return qy_d, qx_d, _inverse_filter(be, lam, alpha), _inverse_filter(be, lam, 0.0)
277
+
278
+
279
+ @functools.lru_cache(maxsize=4)
280
+ def _dct_plan(
281
+ be: Backend, shape: tuple[int, int], pitch: tuple[float, float], regularization: float
282
+ ) -> tuple[Any, Any, Any, Any]:
283
+ """Cosine-series wavenumbers ``pi p / (n pitch)``; regularized and exact inverse ``-lap``."""
284
+ wy = np.pi * np.arange(shape[0]) / (shape[0] * pitch[0])
285
+ wx = np.pi * np.arange(shape[1]) / (shape[1] * pitch[1])
286
+ alpha = regularization * _lam_min([wy**2, wx**2]) ** 2
287
+ lam = wy[:, None] ** 2 + wx[None, :] ** 2
288
+ xp, dt = be.xp, be.real_dtype
289
+ wy_d, wx_d = xp.asarray(wy[:, None], dtype=dt), xp.asarray(wx[None, :], dtype=dt)
290
+ return wy_d, wx_d, _inverse_filter(be, lam, alpha), _inverse_filter(be, lam, 0.0)
291
+
292
+
293
+ @functools.lru_cache(maxsize=4)
294
+ def _fd_inverse(be: Backend, shape: tuple[int, int], pitch: tuple[float, float]) -> Any:
295
+ """Inverse eigenvalues of the Neumann five-point ``-laplacian`` (DCT-II basis)."""
296
+ ey = (2.0 - 2.0 * np.cos(np.pi * np.arange(shape[0]) / shape[0])) / pitch[0] ** 2
297
+ ex = (2.0 - 2.0 * np.cos(np.pi * np.arange(shape[1]) / shape[1])) / pitch[1] ** 2
298
+ return _inverse_filter(be, ey[:, None] + ex[None, :], 0.0)
299
+
300
+
301
+ # ------------------------------------------------------------------ solvers
302
+ def _solve_fft(
303
+ be: Backend,
304
+ b: Any,
305
+ intensity: Any,
306
+ support: Any,
307
+ pitch: tuple[float, float],
308
+ regularization: float,
309
+ ) -> Any:
310
+ """Spectral periodic solve of ``-div(I grad phi) = b`` (Teague) or ``-lap phi = b``.
311
+
312
+ ``intensity`` is the clipped in-focus intensity, or None for ``-lap phi = b``.
313
+ """
314
+ xp = be.xp
315
+ shape = b.shape
316
+ qy, qx, inv, inv_exact = _fft_plan(be, shape, pitch, regularization)
317
+ psi_hat = be.rfft2(b) * inv
318
+ if intensity is None:
319
+ return be.irfft2(psi_hat, shape)
320
+ inv_i = xp.where(support, 1.0 / intensity, 0.0)
321
+ gy = be.irfft2(1j * qy * psi_hat, shape) * inv_i
322
+ gx = be.irfft2(1j * qx * psi_hat, shape) * inv_i
323
+ neg_div = -1j * (qy * be.rfft2(gy) + qx * be.rfft2(gx))
324
+ # The second inversion is bounded (order zero overall): no extra regularization.
325
+ return be.irfft2(neg_div * inv_exact, shape)
326
+
327
+
328
+ def _trig(be: Backend, kind: str, a: Any, axis: int | None, inverse: bool = False) -> Any:
329
+ """Orthonormal type-II DCT (``kind="c"``) or DST (``"s"``) along one axis (None: both)."""
330
+ name = ("i" if inverse else "") + ("dct" if kind == "c" else "dst")
331
+ if be.is_gpu:
332
+ from cupyx.scipy import fft as cfft
333
+
334
+ if axis is None:
335
+ return getattr(cfft, name + "n")(a, type=2, norm="ortho")
336
+ return getattr(cfft, name)(a, type=2, norm="ortho", axis=axis)
337
+ from scipy import fft
338
+
339
+ # Threads only pay off on large grids; on small ones they stall on a busy machine.
340
+ workers = _cpu_workers(a.size)
341
+ if axis is None:
342
+ return getattr(fft, name + "n")(a, type=2, norm="ortho", workers=workers)
343
+ return getattr(fft, name)(a, type=2, norm="ortho", axis=axis, workers=workers)
344
+
345
+
346
+ def _solve_dct(
347
+ be: Backend,
348
+ b: Any,
349
+ intensity: Any,
350
+ support: Any,
351
+ pitch: tuple[float, float],
352
+ regularization: float,
353
+ ) -> Any:
354
+ """Spectral Neumann solve of ``-div(I grad phi) = b`` (Teague) or ``-lap phi = b``.
355
+
356
+ The field is expanded in the cosine series of its even extension (type-II
357
+ DCT), whose derivatives are sine series (type-II DST): differentiation is
358
+ exact for that extension and every gradient vanishes on the boundary.
359
+ """
360
+ xp = be.xp
361
+ wy, wx, inv, inv_exact = _dct_plan(be, b.shape, pitch, regularization)
362
+ psi_hat = _trig(be, "c", b, None) * inv
363
+ if intensity is None:
364
+ return _trig(be, "c", psi_hat, None, True)
365
+ inv_i = xp.where(support, 1.0 / intensity, 0.0)
366
+ # grad: cosine coefficient p -> sine coefficient p - 1, times -w_p.
367
+ sy = xp.zeros_like(psi_hat)
368
+ sy[:-1, :] = -wy[1:] * psi_hat[1:, :]
369
+ sx = xp.zeros_like(psi_hat)
370
+ sx[:, :-1] = -wx[:, 1:] * psi_hat[:, 1:]
371
+ gy = _trig(be, "c", _trig(be, "s", sy, 0, True), 1, True) * inv_i
372
+ gx = _trig(be, "s", _trig(be, "c", sx, 0, True), 1, True) * inv_i
373
+ # div: sine coefficient p - 1 -> cosine coefficient p, times +w_p.
374
+ hy = _trig(be, "c", _trig(be, "s", gy, 0), 1)
375
+ hx = _trig(be, "s", _trig(be, "c", gx, 0), 1)
376
+ div = xp.zeros_like(psi_hat)
377
+ div[1:, :] += wy[1:] * hy[:-1, :]
378
+ div[:, 1:] += wx[:, 1:] * hx[:, :-1]
379
+ return _trig(be, "c", -div * inv_exact, None, True)
380
+
381
+
382
+ def _fd_operator(be: Backend, wy: Any, wx: Any) -> Any:
383
+ """``phi -> D^T diag(w) D phi`` for forward differences ``D`` and edge weights ``w``.
384
+
385
+ ``wy``/``wx`` already include the ``1 / pitch**2`` factors.
386
+ """
387
+
388
+ def apply(phi: Any) -> Any:
389
+ fy = phi[1:, :] - phi[:-1, :]
390
+ fy *= wy
391
+ fx = phi[:, 1:] - phi[:, :-1]
392
+ fx *= wx
393
+ out = be.zeros(phi.shape)
394
+ out[:-1, :] -= fy
395
+ out[1:, :] += fy
396
+ out[:, :-1] -= fx
397
+ out[:, 1:] += fx
398
+ return out
399
+
400
+ return apply
401
+
402
+
403
+ def _solve_pcg(
404
+ be: Backend,
405
+ b: Any,
406
+ intensity: Any,
407
+ mask: Any,
408
+ pitch: tuple[float, float],
409
+ tol: float,
410
+ max_iter: int,
411
+ check_every: int,
412
+ project: Any,
413
+ ) -> tuple[Any, list[float], int, bool, str]:
414
+ """Masked conjugate gradients for ``D^T (I D phi) = b``, DCT-preconditioned.
415
+
416
+ ``D`` is the forward difference; an edge between two support pixels
417
+ carries the mean of their intensities, any other edge carries nothing.
418
+ ``project`` removes the per-region mean (the operator's null space).
419
+ """
420
+ xp = be.xp
421
+ mi = xp.where(mask, intensity, 0.0).astype(be.real_dtype)
422
+ wy = (0.5 / pitch[0] ** 2) * (mi[1:, :] + mi[:-1, :]) * (mask[1:, :] & mask[:-1, :])
423
+ wx = (0.5 / pitch[1] ** 2) * (mi[:, 1:] + mi[:, :-1]) * (mask[:, 1:] & mask[:, :-1])
424
+ i_mean = float(xp.sum(mi)) / max(float(xp.sum(mask)), 1.0)
425
+ inv = _fd_inverse(be, b.shape, pitch) / i_mean
426
+ apply_a = _fd_operator(be, wy.astype(be.real_dtype), wx.astype(be.real_dtype))
427
+
428
+ def precond(r: Any) -> Any:
429
+ # Projecting out the null space keeps single precision stable.
430
+ return project(_trig(be, "c", _trig(be, "c", r, None) * inv, None, True))
431
+
432
+ phi = be.zeros(b.shape)
433
+ history: list[float] = []
434
+ b_norm = math.sqrt(be.dot(b, b))
435
+ if b_norm == 0.0:
436
+ return phi, history, 0, True, "zero right-hand side"
437
+ r = b.copy()
438
+ z = precond(r)
439
+ p = z.copy()
440
+ rz = be.dot(r, z)
441
+ converged = False
442
+ message = f"reached max_iter={max_iter}"
443
+ n_iter = 0
444
+ while n_iter < max_iter:
445
+ n_iter += 1
446
+ ap = apply_a(p)
447
+ pap = be.dot(p, ap)
448
+ if not pap > 0:
449
+ # Only round-off is left (the operator is positive semi-definite).
450
+ rel = math.sqrt(be.dot(r, r)) / b_norm
451
+ history.append(rel)
452
+ converged = rel <= tol
453
+ message = f"stagnated at relative residual {rel:.2e} (round-off)"
454
+ break
455
+ alpha = rz / pap
456
+ phi += alpha * p
457
+ r -= alpha * ap
458
+ if n_iter % check_every == 0 or n_iter == max_iter:
459
+ rel = math.sqrt(be.dot(r, r)) / b_norm
460
+ history.append(rel)
461
+ if rel <= tol:
462
+ converged = True
463
+ message = f"relative residual {rel:.2e} <= tol"
464
+ break
465
+ z = precond(r)
466
+ rz_new = be.dot(r, z)
467
+ if not rz_new > 0:
468
+ rel = math.sqrt(be.dot(r, r)) / b_norm
469
+ history.append(rel)
470
+ converged = rel <= tol
471
+ message = f"preconditioned residual vanished at relative residual {rel:.2e}"
472
+ break
473
+ p *= rz_new / rz
474
+ p += z
475
+ rz = rz_new
476
+ return phi, history, n_iter, converged, message
477
+
478
+
479
+ # -------------------------------------------------------------------- helpers
480
+ def _resolve_backend(array: Any, device: BackendLike, precision: str | None) -> Backend:
481
+ if device is None:
482
+ return backend_of(array, precision)
483
+ return get_backend(device, precision)
484
+
485
+
486
+ def _pitch_pair(pitch: Any) -> tuple[float, float]:
487
+ values = np.broadcast_to(np.asarray(pitch, dtype=np.float64), (2,))
488
+ if not np.all(values > 0) or not np.all(np.isfinite(values)):
489
+ raise ValueError(f"pitch must be positive and finite (metres), got {pitch!r}")
490
+ return float(values[0]), float(values[1])
491
+
492
+
493
+ def _region_labels(be: Backend, mask: Any, per_region: bool) -> tuple[Any, int]:
494
+ """Flattened region labels of ``mask`` (0 outside) and the number of regions."""
495
+ if not per_region:
496
+ return mask.ravel().astype(np.int64), 1
497
+ from scipy import ndimage
498
+
499
+ labels, n_regions = ndimage.label(be.to_numpy(mask))
500
+ return be.xp.asarray(labels.ravel().astype(np.int64)), int(n_regions)
501
+
502
+
503
+ def _remove_region_means(be: Backend, a: Any, labels: Any, n_regions: int) -> Any:
504
+ """Subtract from ``a`` its mean over each labelled region; zero outside them."""
505
+ xp = be.xp
506
+ if n_regions == 1:
507
+ inside = labels.reshape(a.shape) > 0
508
+ mean = xp.sum(xp.where(inside, a, 0.0)) / xp.maximum(xp.sum(inside), 1)
509
+ return xp.where(inside, a - mean, 0.0).astype(a.dtype, copy=False)
510
+ flat = a.ravel()
511
+ sums = xp.bincount(labels, weights=flat, minlength=n_regions + 1)
512
+ counts = xp.bincount(labels, minlength=n_regions + 1)
513
+ means = sums / xp.maximum(counts, 1)
514
+ out = xp.where(labels > 0, flat - means[labels].astype(a.dtype), 0.0)
515
+ return out.reshape(a.shape).astype(a.dtype, copy=False)
516
+
517
+
518
+ # ------------------------------------------------------------------ public
519
+ def tie(
520
+ intensities: Any,
521
+ distances: Sequence[float] | Any,
522
+ *,
523
+ pitch: float | tuple[float, float],
524
+ wavelength: float,
525
+ method: TIEMethod = "dct",
526
+ uniform: bool = False,
527
+ derivative: TIEDerivative = "central",
528
+ order: int | None = None,
529
+ in_focus: Any = None,
530
+ regularization: float = 1e-6,
531
+ intensity_floor: float = 1e-3,
532
+ mask: Any = None,
533
+ tol: float = 1e-6,
534
+ max_iter: int = 2000,
535
+ check_every: int = 10,
536
+ device: BackendLike = None,
537
+ precision: str | None = None,
538
+ ) -> TIEResult:
539
+ """Retrieve the phase at ``z = 0`` from intensities at known defocus.
540
+
541
+ Solves ``-k dI/dz = div(I grad phi)``; see the module notes and the
542
+ transport-of-intensity guide page.
543
+
544
+ Parameters
545
+ ----------
546
+ intensities:
547
+ ``(N, ny, nx)`` intensity images (NumPy or CuPy), any consistent
548
+ units, ``N >= 2``.
549
+ distances:
550
+ ``(N,)`` defocus of each image in metres, relative to the plane where
551
+ the phase is wanted (positive downstream). May include ``0``.
552
+ pitch:
553
+ Pixel pitch in metres, scalar or ``(y, x)``.
554
+ wavelength:
555
+ Wavelength in metres.
556
+ method:
557
+ ``"dct"`` (Neumann boundaries, default), ``"fft"`` (periodic) or
558
+ ``"pcg"`` (exact non-uniform solve inside ``mask``, Neumann).
559
+ uniform:
560
+ Use the uniform-intensity approximation ``I = mean(I0)`` (one Poisson
561
+ solve) instead of the non-uniform solution.
562
+ derivative:
563
+ ``"central"`` (nearest planes about ``z = 0``) or ``"polyfit"``
564
+ (least-squares polynomial over all planes).
565
+ order:
566
+ Polynomial degree for ``"polyfit"``; default ``min(N - 1, 3)``.
567
+ in_focus:
568
+ Optional ``(ny, nx)`` in-focus intensity; overrides the measured
569
+ ``z = 0`` plane and the estimate.
570
+ regularization:
571
+ Dimensionless Tikhonov weight for ``"fft"``/``"dct"``: the lowest
572
+ non-zero frequency is attenuated by ``1 / (1 + regularization)``.
573
+ Ignored by ``"pcg"``.
574
+ intensity_floor:
575
+ Fraction of ``max(I0)`` below which the intensity is clipped: when
576
+ dividing by it (``"fft"``/``"dct"``) and in the ``"pcg"`` weights,
577
+ whose condition number it bounds. The default ``"pcg"`` mask is where
578
+ ``I0`` exceeds it.
579
+ mask:
580
+ Optional ``(ny, nx)`` boolean support. The phase is zero outside it
581
+ and its piston is removed over it. For ``"pcg"`` the equation is
582
+ solved only inside it.
583
+ tol, max_iter, check_every:
584
+ ``"pcg"`` stopping rule: relative residual, iteration limit, and how
585
+ often the residual is checked (each check syncs a GPU).
586
+ device, precision:
587
+ Backend; defaults to the backend of ``intensities``.
588
+
589
+ Returns
590
+ -------
591
+ TIEResult
592
+ Phase (radians) and OPD (metres) on the backend.
593
+ """
594
+ t0 = time.perf_counter()
595
+ be = _resolve_backend(intensities, device, precision)
596
+ xp = be.xp
597
+ stack = be.asarray(intensities, dtype="real")
598
+ if stack.ndim != 3 or stack.shape[0] < 2:
599
+ raise ValueError(
600
+ f"intensities must be an (N, ny, nx) stack with N >= 2, got shape {tuple(stack.shape)}"
601
+ )
602
+ z = np.asarray(distances, dtype=np.float64).ravel()
603
+ if z.size != stack.shape[0]:
604
+ raise ValueError(f"got {z.size} distances for {stack.shape[0]} images")
605
+ if not np.all(np.isfinite(z)) or np.unique(z).size != z.size:
606
+ raise ValueError("distances must be finite and distinct")
607
+ if wavelength <= 0:
608
+ raise ValueError("wavelength must be positive (metres)")
609
+ pitch2 = _pitch_pair(pitch)
610
+ method_name = str(method).lower()
611
+ if method_name not in ("dct", "fft", "pcg"):
612
+ raise ValueError(f"method must be 'dct', 'fft' or 'pcg', got {method!r}")
613
+ if regularization < 0:
614
+ raise ValueError("regularization must be non-negative")
615
+ if not 0 < intensity_floor < 1:
616
+ raise ValueError("intensity_floor must lie in (0, 1)")
617
+
618
+ w_slope, w_value, focus = _axial_weights(z, derivative, order)
619
+ rdt = be.real_dtype
620
+ didz = xp.tensordot(xp.asarray(w_slope, dtype=rdt), stack, axes=1)
621
+ if in_focus is not None:
622
+ i0 = be.asarray(in_focus, dtype="real")
623
+ if i0.shape != stack.shape[1:]:
624
+ raise ValueError(f"in_focus must have shape {stack.shape[1:]}, got {i0.shape}")
625
+ elif focus is not None:
626
+ i0 = stack[focus]
627
+ else:
628
+ i0 = xp.tensordot(xp.asarray(w_value, dtype=rdt), stack, axes=1)
629
+ i_max = float(xp.max(i0))
630
+ if not i_max > 0:
631
+ raise ValueError("the in-focus intensity is not positive anywhere")
632
+ floor = intensity_floor * i_max
633
+
634
+ if mask is not None:
635
+ support = be.asarray(mask).astype(bool)
636
+ if support.shape != i0.shape:
637
+ raise ValueError(f"mask must have shape {i0.shape}, got {support.shape}")
638
+ elif method_name == "pcg":
639
+ support = i0 > floor
640
+ else:
641
+ support = xp.ones(i0.shape, dtype=bool)
642
+ n_support = float(xp.sum(support))
643
+ if n_support == 0:
644
+ raise ValueError("the mask is empty")
645
+ i_mean = float(xp.sum(xp.where(support, i0, 0.0))) / n_support
646
+
647
+ k = 2.0 * math.pi / wavelength
648
+ b = (k * didz).astype(rdt, copy=False)
649
+ history: list[float] = []
650
+ n_iter, converged, message = 0, True, "direct solve"
651
+ labels, n_regions = _region_labels(be, support, method_name == "pcg")
652
+ if method_name == "pcg":
653
+ # Clipping bounds the condition number by 1 / intensity_floor.
654
+ weight = xp.full(i0.shape, i_mean, dtype=rdt) if uniform else xp.maximum(i0, floor)
655
+ # The support's net outflow is unobservable: keep the system consistent.
656
+ b = _remove_region_means(be, b, labels, n_regions)
657
+ phi, history, n_iter, converged, message = _solve_pcg(
658
+ be,
659
+ b,
660
+ weight,
661
+ support,
662
+ pitch2,
663
+ tol,
664
+ int(max_iter),
665
+ max(1, int(check_every)),
666
+ lambda a: _remove_region_means(be, a, labels, n_regions),
667
+ )
668
+ else:
669
+ if uniform:
670
+ rhs, clipped = b / i_mean, None
671
+ else:
672
+ rhs, clipped = b, xp.maximum(i0, floor).astype(rdt, copy=False)
673
+ solver = _solve_fft if method_name == "fft" else _solve_dct
674
+ phi = solver(be, rhs, clipped, support, pitch2, float(regularization))
675
+ phi = _remove_region_means(be, phi.astype(rdt, copy=False), labels, n_regions)
676
+ be.synchronize()
677
+ return TIEResult(
678
+ phase=phi,
679
+ opd=phi * (wavelength / (2.0 * math.pi)),
680
+ intensity=i0,
681
+ didz=didz,
682
+ mask=support,
683
+ wavelength=float(wavelength),
684
+ method=method_name,
685
+ derivative=str(derivative),
686
+ uniform=bool(uniform),
687
+ history=history,
688
+ n_iter=n_iter,
689
+ converged=converged,
690
+ message=message,
691
+ elapsed=time.perf_counter() - t0,
692
+ device=be.device,
693
+ )