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