Phase retrieval
Phase retrieval is a classical inverse problem in optics: recover the phase of a wavefront from intensity measurements alone. In this tutorial we implement the Gerchberg–Saxton algorithm.
The phase problem
Detectors record only intensity, , and discard the phase . Yet the phase is what matters for many applications:
- beam shaping
- adaptive optics
- holography
- optical tweezers
The Gerchberg–Saxton algorithm
Idea
If the intensity distributions in two conjugate planes are known (for example the object plane and the focal plane), the phase can be recovered iteratively.
Steps
- Start from a random phase.
- Impose the known amplitude in plane 1.
- Transform to plane 2 (FFT or physical propagation).
- Impose the known amplitude in plane 2.
- Transform back.
- Repeat until convergence.
Step 1: Setup
import torch
import matplotlib.pyplot as plt
from svetlanna import SimulationParameters, Wavefront
from svetlanna.elements import ThinLens, FreeSpace
from svetlanna.units import ureg
params = SimulationParameters.from_ranges(
x_range=(-2*ureg.mm, 2*ureg.mm), x_points=256,
y_range=(-2*ureg.mm, 2*ureg.mm), y_points=256,
wavelength=632.8*ureg.nm,
)
X, Y = params.meshgrid(x_axis="x", y_axis="y")Step 2: Test data
Two “known” amplitude distributions:
def create_letter_A(X, Y, size=1*ureg.mm):
"""Build a mask shaped like the letter A."""
left = (X > -size*0.4) & (X < -size*0.3) & (Y > -size*0.5) & (Y < size*0.5)
right = (X > size*0.3) & (X < size*0.4) & (Y > -size*0.5) & (Y < size*0.5)
bar = (X > -size*0.3) & (X < size*0.3) & (Y > -size*0.1) & (Y < size*0.1)
top = (Y > size*0.3) & (Y < size*0.5) & (torch.abs(X) < (size*0.5 - Y)*0.8)
return (left | right | bar | top).float()
source_amplitude = create_letter_A(X, Y)
# target in the focal plane — an annulus
r = torch.sqrt(X**2 + Y**2)
target_amplitude = ((r > 0.3*ureg.mm) & (r < 0.5*ureg.mm)).float()
source_amplitude = source_amplitude / source_amplitude.max()
target_amplitude = target_amplitude / target_amplitude.sum().sqrt()Step 3: Forward and backward transforms
For simplicity we model propagation with an FFT:
def forward_propagate(field):
"""Forward transform (FFT)."""
return torch.fft.fftshift(torch.fft.fft2(torch.fft.ifftshift(field)))
def backward_propagate(field):
"""Inverse transform (IFFT)."""
return torch.fft.fftshift(torch.fft.ifft2(torch.fft.ifftshift(field)))Step 4: The algorithm
def gerchberg_saxton(source_amplitude, target_amplitude, n_iterations=100):
"""
Gerchberg–Saxton phase retrieval.
Parameters
----------
source_amplitude : Tensor
Amplitude in the source plane.
target_amplitude : Tensor
Amplitude in the target plane.
n_iterations : int
Number of iterations.
Returns
-------
phase : Tensor
Recovered phase.
errors : list
Error history.
"""
phase = torch.rand_like(source_amplitude) * 2 * torch.pi
errors = []
for iteration in range(n_iterations):
# 1. field in the source plane with the current phase
source_field = source_amplitude * torch.exp(1j * phase)
# 2. propagate forward
target_field = forward_propagate(source_field)
# 3. measure the error
error = torch.mean((torch.abs(target_field) - target_amplitude)**2)
errors.append(error.item())
# 4. replace the amplitude, keep the phase
target_phase = torch.angle(target_field)
target_field = target_amplitude * torch.exp(1j * target_phase)
# 5. propagate back
source_field = backward_propagate(target_field)
# 6. extract the phase
phase = torch.angle(source_field)
if iteration % 20 == 0:
print(f"Iteration {iteration}: error = {error:.6f}")
return phase, errors
recovered_phase, errors = gerchberg_saxton(
source_amplitude, target_amplitude, n_iterations=100
)Step 5: Checking the result
source_field = source_amplitude * torch.exp(1j * recovered_phase)
target_field = forward_propagate(source_field)
fig, axes = plt.subplots(2, 3, figsize=(15, 10))
axes[0, 0].imshow(source_amplitude.cpu(), cmap='gray')
axes[0, 0].set_title('Source amplitude')
axes[0, 1].imshow(recovered_phase.cpu(), cmap='twilight')
axes[0, 1].set_title('Recovered phase')
axes[0, 2].plot(errors)
axes[0, 2].set_xlabel('Iteration')
axes[0, 2].set_ylabel('Error')
axes[0, 2].set_title('Convergence')
axes[0, 2].set_yscale('log')
axes[1, 0].imshow(target_amplitude.cpu()**2, cmap='hot')
axes[1, 0].set_title('Target intensity')
axes[1, 1].imshow(torch.abs(target_field).cpu()**2, cmap='hot')
axes[1, 1].set_title('Achieved intensity')
diff = torch.abs(torch.abs(target_field)**2 - target_amplitude**2)
axes[1, 2].imshow(diff.cpu(), cmap='hot')
axes[1, 2].set_title('Difference')
plt.tight_layout()
plt.show()Step 6: Constraining the phase
Real modulators cover a limited phase range:
def gerchberg_saxton_constrained(source_amplitude, target_amplitude,
n_iterations=100, phase_range=(0, 2*torch.pi)):
"""GS with a restricted phase range."""
lo, hi = phase_range
phase = torch.rand_like(source_amplitude) * (hi - lo) + lo
errors = []
for iteration in range(n_iterations):
source_field = source_amplitude * torch.exp(1j * phase)
target_field = forward_propagate(source_field)
error = torch.mean((torch.abs(target_field) - target_amplitude)**2)
errors.append(error.item())
target_phase = torch.angle(target_field)
target_field = target_amplitude * torch.exp(1j * target_phase)
source_field = backward_propagate(target_field)
phase = torch.clamp(torch.angle(source_field), lo, hi)
return phase, errorsStep 7: Using real SVETlANNa elements
Instead of a bare FFT, propagate through an actual lens. Backward propagation uses a negative distance and a negative focal length:
def gerchberg_saxton_optical(params, source_amplitude, target_amplitude,
focal_length, n_iterations=100):
"""GS driven by a physical optical system (lens + free space)."""
lens = ThinLens(params, focal_length=focal_length)
prop = FreeSpace(params, distance=focal_length, method="zpASM")
# backward branch
prop_back = FreeSpace(params, distance=-focal_length, method="zpASM")
lens_inv = ThinLens(params, focal_length=-focal_length)
phase = torch.rand_like(source_amplitude) * 2 * torch.pi
errors = []
for iteration in range(n_iterations):
wf = Wavefront((source_amplitude * torch.exp(1j * phase)).to(torch.complex64))
target_field = prop(lens(wf))
error = torch.mean((torch.sqrt(target_field.intensity) - target_amplitude)**2)
errors.append(error.item())
target_field = Wavefront(
(target_amplitude * torch.exp(1j * target_field.phase)).to(torch.complex64)
)
phase = lens_inv(prop_back(target_field)).phase
return phase, errorsFreeSpace has no reverse() method — build a second FreeSpace with a negative
distance to propagate backwards. ThinLens, DiffractiveLayer and
SpatialLightModulator do provide reverse().
Step 8: Gradient-based retrieval
An alternative is to let PyTorch autograd do the work:
def phase_retrieval_gradient(source_amplitude, target_amplitude,
n_iterations=500, lr=0.1):
"""Phase retrieval by gradient descent."""
phase = torch.nn.Parameter(torch.rand_like(source_amplitude) * 2 * torch.pi)
optimizer = torch.optim.Adam([phase], lr=lr)
errors = []
for iteration in range(n_iterations):
optimizer.zero_grad()
source_field = source_amplitude * torch.exp(1j * phase)
target_field = forward_propagate(source_field)
loss = torch.mean((torch.abs(target_field) - target_amplitude)**2)
loss.backward()
optimizer.step()
errors.append(loss.item())
if iteration % 100 == 0:
print(f"Iteration {iteration}: loss = {loss.item():.6f}")
return phase.detach(), errorsFull code
import torch
from svetlanna import SimulationParameters
from svetlanna.units import ureg
params = SimulationParameters.from_ranges(
x_range=(-2*ureg.mm, 2*ureg.mm), x_points=256,
y_range=(-2*ureg.mm, 2*ureg.mm), y_points=256,
wavelength=632.8*ureg.nm,
)
X, Y = params.meshgrid(x_axis="x", y_axis="y")
# target — an annulus
r = torch.sqrt(X**2 + Y**2)
target = ((r > 0.3*ureg.mm) & (r < 0.5*ureg.mm)).float()
target = target / target.sum().sqrt()
# source — a Gaussian beam
source = torch.exp(-(X**2 + Y**2) / (0.5*ureg.mm)**2)
phase = torch.rand(256, 256) * 2 * torch.pi
for i in range(100):
field = source * torch.exp(1j * phase)
spectrum = torch.fft.fftshift(torch.fft.fft2(torch.fft.ifftshift(field)))
spectrum = target * torch.exp(1j * torch.angle(spectrum))
field = torch.fft.fftshift(torch.fft.ifft2(torch.fft.ifftshift(spectrum)))
phase = torch.angle(field)
print("Retrieval finished.")Conclusions
- Gerchberg–Saxton recovers phase efficiently when both intensities are known.
- Convergence is typically reached in 50–100 iterations.
- The gradient-based variant can be more accurate but costs more compute.
- SVETlANNa lets you replace the FFT with a realistic propagation model.
What next?
- Phase grating on an SLM — displaying the retrieved phase
- Optimisation — gradient-based methods