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 minimiseMatching 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_yfwhm() 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_energyPyTorch optimisers
Adam (recommended)
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 phaseSee 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_lossTotal 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
- Optical systems — building setups
- Tutorial: phase retrieval — the GS algorithm
- Tutorial: phase grating on an SLM — differentiable quantisation