Skip to Content
DocsTutorialsPhase retrieval

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, I=∣E∣2I = |E|^2, and discard the phase ϕ=arg⁡(E)\phi = \arg(E). 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

  1. Start from a random phase.
  2. Impose the known amplitude in plane 1.
  3. Transform to plane 2 (FFT or physical propagation).
  4. Impose the known amplitude in plane 2.
  5. Transform back.
  6. 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, errors

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

FreeSpace 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(), errors

Full 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

  1. Gerchberg–Saxton recovers phase efficiently when both intensities are known.
  2. Convergence is typically reached in 50–100 iterations.
  3. The gradient-based variant can be more accurate but costs more compute.
  4. SVETlANNa lets you replace the FFT with a realistic propagation model.

What next?