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,255 @@
1
+ """Gerchberg-Saxton and its multi-plane (Misell) form for focal-plane data.
2
+
3
+ The pupil amplitude is known, the image intensities are measured, and the
4
+ algorithm alternates projections between the two (Gerchberg & Saxton, Optik
5
+ 35, 237, 1972). With several diversity channels it is Misell's multi-plane
6
+ algorithm (J. Phys. D 6, L6, 1973) in its parallel form: each channel's
7
+ modulus-corrected field is propagated back, its diversity removed, and the
8
+ pupil estimates are averaged as phasors (as in JWST's Hybrid Diversity
9
+ Algorithm, Dean et al. 2006).
10
+
11
+ The modulus projection keeps the model field wherever the data are not
12
+ measured (outside the detector window when the engine is an FFT, or at
13
+ zero-weight pixels), which is the exact projection onto the measurement set.
14
+ The result is a wrapped phase; it is unwrapped (weighted least squares) and
15
+ optionally fitted to a modal basis. GS is fast and has a wide capture range,
16
+ which makes it the standard initializer for nonlinear optimization.
17
+ """
18
+
19
+ from __future__ import annotations
20
+
21
+ import math
22
+ import time
23
+ import warnings
24
+ from collections.abc import Callable
25
+ from typing import Any
26
+
27
+ import numpy as np
28
+
29
+ from ..basis import Basis
30
+ from ..focal import FocalPlaneModel
31
+ from ..propagation import FFTPropagator, FocalPlanePropagator, Propagator
32
+ from ..result import Result
33
+ from ..retrieval import FocalPlaneProblem
34
+ from ..unwrap import unwrap_phase
35
+
36
+ __all__ = ["gerchberg_saxton"]
37
+
38
+
39
+ class _GSPlan:
40
+ """Monochromatic propagation plan used by the iterations."""
41
+
42
+ def __init__(self, model: FocalPlaneModel) -> None:
43
+ be = model.backend
44
+ xp = be.xp
45
+ self.backend = be
46
+ os_ = model.oversample
47
+ my, mx = model.image_shape
48
+ self.window = (my * os_, mx * os_)
49
+ samples = model.wavelength / (model.pupil.pitch * model.pixel_scale / os_)
50
+ n_fft = round(samples)
51
+ offset = (model.offset[0] * os_, model.offset[1] * os_)
52
+ self.full = (
53
+ abs(samples - n_fft) < 1e-9 * samples
54
+ and n_fft >= max(model.pupil.shape)
55
+ and n_fft >= max(self.window)
56
+ )
57
+ self.prop: Propagator
58
+ if self.full:
59
+ # Propagate to the whole FFT grid; locate the detector window inside it.
60
+ starts, full_offset = [], []
61
+ for axis in range(2):
62
+ s = (n_fft - self.window[axis]) // 2
63
+ starts.append(s)
64
+ full_offset.append(offset[axis] + (n_fft - self.window[axis]) / 2.0 - s)
65
+ self.prop = FFTPropagator(
66
+ model.pupil.shape, (n_fft, n_fft), n_fft, offset=tuple(full_offset), backend=be
67
+ )
68
+ self.slices = (
69
+ slice(starts[0], starts[0] + self.window[0]),
70
+ slice(starts[1], starts[1] + self.window[1]),
71
+ )
72
+ else:
73
+ self.prop = FocalPlanePropagator(
74
+ model.pupil.shape,
75
+ model.pupil.pitch,
76
+ [model.wavelength],
77
+ model.pixel_scale / os_,
78
+ self.window,
79
+ offset=offset,
80
+ method="mft",
81
+ backend=be,
82
+ )
83
+ self.slices = (slice(None), slice(None))
84
+ self.xp = xp
85
+
86
+ def forward(self, u: Any) -> Any:
87
+ if self.full:
88
+ return self.prop.forward(u)
89
+ return self.prop.forward(u[:, None])[:, 0]
90
+
91
+ def adjoint(self, e: Any) -> Any:
92
+ if self.full:
93
+ return self.prop.adjoint(e)
94
+ return self.prop.adjoint(e[:, None])[:, 0]
95
+
96
+
97
+ def gerchberg_saxton(
98
+ model: FocalPlaneModel,
99
+ images: Any,
100
+ *,
101
+ iterations: int = 200,
102
+ start: Any = None,
103
+ basis: Basis | int | None = None,
104
+ background: Any = 0.0,
105
+ weights: Any = None,
106
+ momentum: float = 0.0,
107
+ tol: float = 1e-6,
108
+ check_every: int = 10,
109
+ unwrap: bool = True,
110
+ callback: Callable[[int, Any, float], bool | None] | None = None,
111
+ ) -> Result:
112
+ """Gerchberg-Saxton / Misell phase retrieval with a known pupil amplitude.
113
+
114
+ Parameters
115
+ ----------
116
+ model:
117
+ Forward model. Broadband models are run at the reference wavelength
118
+ (a warning is issued when the fractional bandwidth exceeds 5%).
119
+ images:
120
+ ``(K, my, mx)`` measured images (any linear units).
121
+ iterations:
122
+ Maximum number of iterations.
123
+ start:
124
+ Initial OPD map in metres (default flat).
125
+ basis:
126
+ Optional basis (or number of Zernike modes) the unwrapped phase is
127
+ fitted to; the returned OPD is then the modal fit.
128
+ background:
129
+ Per-channel background subtracted before taking square roots.
130
+ weights:
131
+ Per-pixel weights; zero marks unmeasured pixels, whose model value is
132
+ kept.
133
+ momentum:
134
+ Over-relaxation ``beta`` of the accelerated GS update
135
+ ``theta <- P(theta) + beta (P(theta) - P(theta_prev))`` (0 = plain GS).
136
+ tol:
137
+ Stop when the relative change of the error metric over
138
+ ``check_every`` iterations is below ``tol``.
139
+ check_every:
140
+ Error-metric interval (each check synchronizes the GPU).
141
+ unwrap:
142
+ Unwrap the final phase (otherwise the OPD is the wrapped phase).
143
+ callback:
144
+ ``callback(iteration, phase, error)``; return True to stop.
145
+
146
+ Returns
147
+ -------
148
+ Result
149
+ ``extra["wrapped_phase"]`` holds the wrapped pupil phase and
150
+ ``history`` the normalized amplitude error
151
+ ``sum (|E| - sqrt(d))^2 / sum d`` over measured pixels.
152
+ """
153
+ t0 = time.perf_counter()
154
+ be = model.backend
155
+ xp = be.xp
156
+ bandwidth = (model.wavelengths.max() - model.wavelengths.min()) / model.wavelength
157
+ if bandwidth > 0.05:
158
+ warnings.warn(
159
+ f"Gerchberg-Saxton runs monochromatically at the reference wavelength; this "
160
+ f"model has {bandwidth:.0%} bandwidth. Refine with solvephase.solve().",
161
+ stacklevel=2,
162
+ )
163
+ plan = _GSPlan(model)
164
+ data = be.asarray(images, dtype="real")
165
+ if data.ndim == 2:
166
+ data = data[None]
167
+ k = model.n_channels
168
+ if data.shape != (k, *model.image_shape):
169
+ raise ValueError(
170
+ f"images have shape {tuple(data.shape)}, expected {(k, *model.image_shape)}"
171
+ )
172
+ bg = be.asarray(np.broadcast_to(np.asarray(background, dtype=np.float64), (k,)), dtype="real")
173
+ signal = xp.maximum(data - bg[:, None, None], 0.0)
174
+ os_ = model.oversample
175
+ if os_ > 1:
176
+ signal = xp.repeat(xp.repeat(signal, os_, axis=-2), os_, axis=-1) / os_**2
177
+ measured = xp.ones(signal.shape, dtype=bool)
178
+ if weights is not None:
179
+ w = be.asarray(np.broadcast_to(be.to_numpy(weights), tuple(data.shape)), dtype="real") > 0
180
+ if os_ > 1:
181
+ w = xp.repeat(xp.repeat(w, os_, axis=-2), os_, axis=-1)
182
+ measured = w
183
+ sqrt_signal = xp.sqrt(signal)
184
+ sig_energy = xp.sum(signal * measured, axis=(1, 2))
185
+
186
+ amp = be.asarray(model.pupil.amplitude, dtype="real")
187
+ div = model._div # (K, ny, nx) radians at the reference wavelength
188
+ if start is None:
189
+ theta = xp.zeros(model.pupil.shape, dtype=be.real_dtype)
190
+ elif isinstance(start, Result):
191
+ theta = be.asarray(start.phase, dtype="real")
192
+ else:
193
+ theta = be.asarray(start, dtype="real") * (2 * math.pi / model.wavelength)
194
+ ys, xs = plan.slices
195
+ prev_proj = None
196
+ history: list[float] = []
197
+ times: list[float] = []
198
+ err_prev = math.inf
199
+ message, converged = "iteration limit reached", False
200
+ it = 0
201
+ for it in range(1, iterations + 1):
202
+ u = amp * xp.exp(1j * (theta[None] + div))
203
+ u = u.astype(be.complex_dtype, copy=False)
204
+ e = plan.forward(u)
205
+ win = e[:, ys, xs]
206
+ mag = xp.abs(win)
207
+ # Match the data's energy to the model's in the measured window.
208
+ model_energy = xp.sum((mag * mag) * measured, axis=(1, 2))
209
+ scale = xp.sqrt(model_energy / xp.maximum(sig_energy, 1e-300))
210
+ target = sqrt_signal * scale[:, None, None]
211
+ new_win = xp.where(measured, target * win / xp.maximum(mag, 1e-30), win)
212
+ check = it % check_every == 0 or it == iterations
213
+ if check:
214
+ resid = xp.sum(((mag - target) ** 2) * measured)
215
+ err = float(resid / xp.maximum(xp.sum(model_energy), 1e-300))
216
+ history.append(err)
217
+ times.append(time.perf_counter() - t0)
218
+ e[:, ys, xs] = new_win
219
+ back = plan.adjoint(e) * xp.exp(-1j * div)
220
+ proj = xp.angle(xp.sum(back, axis=0))
221
+ if momentum and prev_proj is not None:
222
+ step = xp.angle(xp.exp(1j * (proj - prev_proj)))
223
+ theta = proj + momentum * step
224
+ else:
225
+ theta = proj
226
+ prev_proj = proj
227
+ if check:
228
+ if callback is not None and callback(it, theta, err):
229
+ message, converged = "stopped by callback", True
230
+ break
231
+ if abs(err_prev - err) <= tol * max(err, 1e-300):
232
+ message, converged = "error metric stalled below tol", True
233
+ break
234
+ err_prev = err
235
+ theta = xp.where(be.asarray(model.pupil.mask), theta, 0.0)
236
+ wrapped = theta
237
+ if unwrap:
238
+ theta = unwrap_phase(theta, model.pupil.mask, weights=model.pupil.amplitude, device=be)
239
+ theta = be.asarray(theta, dtype="real")
240
+ if isinstance(basis, (int, np.integer)):
241
+ basis = Basis.zernike(model.pupil, int(basis))
242
+ opd = theta * (model.wavelength / (2 * math.pi))
243
+ problem = FocalPlaneProblem(model, data, basis=basis, background=background)
244
+ x = problem.initial(be.to_numpy(opd))
245
+ result = problem.result(
246
+ x, method="gs" if k == 1 else "misell", elapsed=time.perf_counter() - t0
247
+ )
248
+ result.history, result.times = history, times
249
+ result.n_iter, result.converged, result.message = it, converged, message
250
+ result.loss = history[-1] if history else math.nan
251
+ result.extra["wrapped_phase"] = wrapped
252
+ if basis is None:
253
+ # Keep the unwrapped zonal phase exactly (the zonal fit is the identity).
254
+ result.opd, result.phase = opd, theta
255
+ return result