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,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
|
+
)
|