Skip to Content
DocsGuidesOptimisation and training

Optimisation and training

SVETlANNa is fully differentiable thanks to its PyTorch backend, so optical systems can be designed by gradient descent rather than by hand.

Automatic differentiation basics

Trainable parameters

Use nn.Parameter — or SVETlANNa’s Parameter / ConstrainedParameter — for anything you want to optimise:

import torch import torch.nn as nn from svetlanna import ConstrainedParameter from svetlanna.elements import ThinLens, FreeSpace from svetlanna.units import ureg # a lens whose focal length is trainable but physically bounded lens = ThinLens( params, focal_length=ConstrainedParameter( 100e-3, min_value=50e-3, max_value=200e-3 ), ) # a trainable phase mask phase_mask = nn.Parameter(torch.rand(512, 512) * 2 * torch.pi)

ConstrainedParameter reparametrises the value through a sigmoid, so the optimiser is unconstrained while the physical value never leaves the interval.

Computing gradients

wf_out = system(wf_in) loss = loss_function(wf_out) loss.backward() for name, param in system.named_parameters(): print(f"{name}: grad norm = {param.grad.norm():.3e}")

Loss functions

Maximising intensity at a point

def peak_intensity_loss(wf, target_x, target_y, params): """Maximise the intensity at a given point.""" ix = torch.argmin(torch.abs(params.x - target_x)) iy = torch.argmin(torch.abs(params.y - target_y)) return -wf.intensity[iy, ix] # negated: we minimise

Matching a target distribution

def target_distribution_loss(wf, target_intensity): """MSE between the normalised intensity and a target.""" return nn.functional.mse_loss( wf.intensity / wf.intensity.max(), target_intensity / target_intensity.max() )

Minimising the spot size

def fwhm_loss(wf, params): fwhm_x, fwhm_y = wf.fwhm(params) return fwhm_x + fwhm_y

fwhm() involves a comparison and an index difference, so it carries no useful gradient. Use it as a metric, not as a loss — minimise a smooth surrogate such as the second moment of the intensity instead.

Energy in a target region

def efficiency_loss(wf, target_mask): """Maximise the fraction of energy inside a region.""" total_energy = wf.intensity.sum() target_energy = (wf.intensity * target_mask).sum() return -target_energy / total_energy

PyTorch optimisers

optimizer = torch.optim.Adam(system.parameters(), lr=1e-3)

SGD with momentum

optimizer = torch.optim.SGD(system.parameters(), lr=1e-2, momentum=0.9)

LBFGS (for fine convergence)

optimizer = torch.optim.LBFGS(system.parameters(), lr=1.0, max_iter=20) def closure(): optimizer.zero_grad() wf_out = system(wf_in) loss = loss_function(wf_out) loss.backward() return loss optimizer.step(closure)

The training loop

Basic loop

import torch from svetlanna import SimulationParameters, Wavefront from svetlanna.units import ureg params = SimulationParameters.from_ranges( x_range=(-1*ureg.mm, 1*ureg.mm), x_points=256, y_range=(-1*ureg.mm, 1*ureg.mm), y_points=256, wavelength=632.8*ureg.nm ) wf_in = Wavefront.gaussian_beam(params, waist_radius=0.3*ureg.mm) system = OptimizableSystem(params) optimizer = torch.optim.Adam(system.parameters(), lr=1e-3) losses = [] for epoch in range(1000): optimizer.zero_grad() wf_out = system(wf_in) loss = loss_function(wf_out) loss.backward() optimizer.step() losses.append(loss.item()) if epoch % 100 == 0: print(f"Epoch {epoch}: loss = {loss.item():.6f}")

With validation

for epoch in range(1000): system.train() optimizer.zero_grad() train_loss = loss_function(system(wf_in)) train_loss.backward() optimizer.step() system.eval() with torch.no_grad(): val_loss = loss_function(system(wf_validation)) print(f"Epoch {epoch}: train={train_loss:.6f}, val={val_loss:.6f}")

Worked examples

Optimising an SLM phase mask

import torch import torch.nn as nn from svetlanna.elements import ThinLens, FreeSpace from svetlanna.units import ureg class SLMOptimization(nn.Module): def __init__(self, params, focal_length, grid=256): super().__init__() self.phase = nn.Parameter(torch.zeros(grid, grid)) self.lens = ThinLens(params, focal_length=focal_length) self.prop = FreeSpace(params, distance=focal_length, method="zpASM") def forward(self, wf): # keep the phase inside [0, 2π] phase = torch.sigmoid(self.phase) * 2 * torch.pi wf = wf * torch.exp(1j * phase) return self.prop(self.lens(wf)) X, Y = params.meshgrid(x_axis='x', y_axis='y') r = torch.sqrt(X**2 + Y**2) target = ((r > 10*ureg.um) & (r < 20*ureg.um)).float() system = SLMOptimization(params, focal_length=100*ureg.mm) optimizer = torch.optim.Adam(system.parameters(), lr=0.1) for epoch in range(500): optimizer.zero_grad() wf_out = system(wf_in) loss = nn.functional.mse_loss( wf_out.intensity / wf_out.intensity.max(), target ) loss.backward() optimizer.step()

Gerchberg–Saxton phase retrieval

def gerchberg_saxton(source_amplitude, target_amplitude, n_iterations=100): """Iterative phase retrieval.""" phase = torch.rand_like(source_amplitude) * 2 * torch.pi for _ in range(n_iterations): field = source_amplitude * torch.exp(1j * phase) spectrum = torch.fft.fft2(field) spectrum = target_amplitude * torch.exp(1j * torch.angle(spectrum)) field = torch.fft.ifft2(spectrum) phase = torch.angle(field) return phase

See the phase retrieval tutorial for the complete version.

GPU acceleration

device = "cuda" if torch.cuda.is_available() else "cpu" params.to(device) # in-place wf_in = wf_in.to(device) system = system.to(device) target = target.to(device)

Mixed precision (torch.autocast) is of limited use here: complex-valued FFTs are not covered by float16 autocasting, and reduced precision quickly degrades phase accuracy. Prefer float32 and reduce the grid size if memory is tight.

Regularisation

L2 penalty on the phase

def regularized_loss(wf_out, target, phase, lambda_reg=0.01): main_loss = nn.functional.mse_loss(wf_out.intensity, target) reg_loss = lambda_reg * (phase**2).mean() return main_loss + reg_loss

Total variation (phase smoothness)

def total_variation(phase): """Encourage a smooth, manufacturable phase profile.""" diff_x = torch.abs(phase[:, 1:] - phase[:, :-1]) diff_y = torch.abs(phase[1:, :] - phase[:-1, :]) return diff_x.mean() + diff_y.mean()

Bounding the phase range

# soft bound through a sigmoid — keeps gradients everywhere phase_constrained = torch.sigmoid(self.phase) * 2 * torch.pi # hard bound — zero gradient outside the interval phase_constrained = torch.clamp(self.phase, 0, 2 * torch.pi)

Visualising training

import matplotlib.pyplot as plt fig, axes = plt.subplots(1, 3, figsize=(15, 4)) axes[0].plot(losses) axes[0].set_xlabel('Epoch') axes[0].set_ylabel('Loss') axes[0].set_title('Training loss') axes[0].set_yscale('log') axes[1].imshow(wf_out.intensity.detach().cpu(), cmap='hot') axes[1].set_title('Output intensity') axes[2].imshow(system.phase.detach().cpu(), cmap='twilight') axes[2].set_title('Optimised phase') plt.tight_layout() plt.show()

See also