Skip to Content
DocsGuidesTransforms

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 [0,1][0, 1]. 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 modulation — image values become phase, the amplitude stays constant:

E(x,y)=exp⁡(i[2πI(x,y)−π])E(x, y) = \exp\left(i\left[2\pi I(x, y) - \pi\right]\right)

transform = ToWavefront(modulation_type='phase') wf = transform(image) # image ∈ [0, 1] # wf.phase ∈ (-π, π) # |wf| = 1

GaussModulation

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

ParameterDescription
sim_paramsThe simulation grid the Gaussian is sampled on
fwhm_x, fwhm_yFull width at half maximum, SI units
peak_x, peak_yCentre 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 [0,1][0, 1] 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