Transforms
The svetlanna.transforms module prepares input data for optical processing.
ToWavefront
Converts an image tensor into a wavefront.
from svetlanna.transforms import ToWavefront
# phase modulation: image values → phase
to_wf_phase = ToWavefront(modulation_type='phase')
# amplitude modulation: image values → amplitude
to_wf_amp = ToWavefront(modulation_type='amp')
# both at once
to_wf_both = ToWavefront(modulation_type='amp&phase')The input must be a tensor of shape [C, H, W] with values in . A single-channel
image has its channel dimension squeezed out; with several channels the channel axis is
kept and interpreted as the wavelength axis.
Modulation types
phase
Phase modulation — image values become phase, the amplitude stays constant:
transform = ToWavefront(modulation_type='phase')
wf = transform(image) # image ∈ [0, 1]
# wf.phase ∈ (-π, π)
# |wf| = 1GaussModulation
Multiplies a wavefront by a Gaussian envelope — a model of illumination by a real laser beam.
from svetlanna.transforms import GaussModulation
from svetlanna.units import ureg
gauss_mod = GaussModulation(
sim_params=params,
fwhm_x=2.0*ureg.mm, # FWHM along X
fwhm_y=2.0*ureg.mm, # FWHM along Y
peak_x=0.0, # centre along X
peak_y=0.0 # centre along Y
)
wf_modulated = gauss_mod(wf)
# wf_modulated = wf * Gaussian(x, y)Parameters
| Parameter | Description |
|---|---|
sim_params | The simulation grid the Gaussian is sampled on |
fwhm_x, fwhm_y | Full width at half maximum, SI units |
peak_x, peak_y | Centre position, SI units |
Composing transforms
Transforms are nn.Modules, so they compose with nn.Sequential or
torchvision.transforms.Compose:
import torch.nn as nn
from svetlanna.transforms import ToWavefront, GaussModulation
from svetlanna.units import ureg
class InputPipeline(nn.Module):
def __init__(self, params):
super().__init__()
self.to_wf = ToWavefront(modulation_type='phase')
self.gauss = GaussModulation(
sim_params=params,
fwhm_x=1.5*ureg.mm,
fwhm_y=1.5*ureg.mm
)
def forward(self, image):
wf = self.to_wf(image)
return self.gauss(wf)
pipeline = InputPipeline(params)
wf = pipeline(image)Normalising the input
Make sure the image is normalised to before applying ToWavefront; the phase
mapping assumes that range.
import torch
def normalize_image(image):
"""Rescale to [0, 1]."""
min_val = image.min()
max_val = image.max()
return (image - min_val) / (max_val - min_val + 1e-8)
image_norm = normalize_image(raw_image)
wf = ToWavefront('phase')(image_norm)Example: an MNIST classification pipeline
import torch
from torch.utils.data import DataLoader
from torchvision import datasets, transforms as T
from svetlanna import SimulationParameters, LinearOpticalSetup
from svetlanna.elements import DiffractiveLayer, FreeSpace
from svetlanna.transforms import ToWavefront, GaussModulation
from svetlanna.detector import Detector, DetectorProcessorClf
from svetlanna.units import ureg
params = SimulationParameters.from_ranges(
x_range=(-5*ureg.mm, 5*ureg.mm), x_points=28, # MNIST resolution
y_range=(-5*ureg.mm, 5*ureg.mm), y_points=28,
wavelength=632.8*ureg.nm
)
class MNISTToOptical(torch.nn.Module):
def __init__(self, params):
super().__init__()
self.to_wf = ToWavefront(modulation_type='phase')
self.gauss = GaussModulation(
sim_params=params, fwhm_x=3*ureg.mm, fwhm_y=3*ureg.mm
)
def forward(self, image):
wf = self.to_wf(image)
return self.gauss(wf)
class DONN(torch.nn.Module):
def __init__(self, params, num_classes, grid=28):
super().__init__()
self.input_transform = MNISTToOptical(params)
self.optical = LinearOpticalSetup([
DiffractiveLayer(
params, mask=torch.nn.Parameter(torch.rand(grid, grid) * 2 * torch.pi)
),
FreeSpace(params, distance=10*ureg.mm, method='zpASM'),
DiffractiveLayer(
params, mask=torch.nn.Parameter(torch.rand(grid, grid) * 2 * torch.pi)
),
FreeSpace(params, distance=10*ureg.mm, method='zpASM'),
])
self.detector = Detector(params, func='intensity')
self.processor = DetectorProcessorClf(
num_classes=num_classes,
simulation_parameters=params,
segmentation_type='strips'
)
def forward(self, images):
wf = self.input_transform(images)
wf = self.optical(wf)
intensity = self.detector(wf)
return self.processor.batch_forward(intensity)
model = DONN(params, num_classes=10)
dataset = datasets.MNIST('./data', train=True, download=True, transform=T.ToTensor())
loader = DataLoader(dataset, batch_size=32)
images, labels = next(iter(loader))
logits = model(images)
print(f"Output shape: {logits.shape}") # (32, 10)Applying ToWavefront per sample inside the Dataset transform — rather than to a whole
batch — keeps the shape convention [C, H, W] and lets DataLoader do the batching.
See also
- Detectors — detection and classification
- Wavefronts — the field object
- Optical systems — building the network