Detectors
The svetlanna.detector module turns an optical field into a measurable quantity and
post-processes it for classification tasks.
Detector
Detector converts a wavefront into an intensity image.
from svetlanna.detector import Detector
detector = Detector(
simulation_parameters=params,
func='intensity' # currently the only mode
)
detector_image = detector(wavefront)
# detector_image == |wavefront|² == wavefront.intensityDetector is equivalent to reading wf.intensity, but as an nn.Module it slots into a
Sequential pipeline and shows up in state_dict().
DetectorProcessorClf
DetectorProcessorClf reads out a classification result by summing the intensity inside
zones of the detector — one zone per class.
from svetlanna.detector import DetectorProcessorClf
processor = DetectorProcessorClf(
num_classes=10,
simulation_parameters=params,
segmentation_type='strips',
segments_zone_size=None, # optional: restrict the zones to a sub-region
device='cpu',
)
# a single image
class_scores = processor(detector_image)
# shape: (1, num_classes)
# a batch
batch_scores = processor.batch_forward(batch_detector_images)
# shape: (batch_size, num_classes)
# the integral over one class zone
integral = processor.batch_zone_integral(batch_images, ind_class=0)
# shape: (batch_size,)Segmentation types
| Type | Description |
|---|---|
'strips' | Vertical strips placed symmetrically about the centre |
Visualising the zones
segmented_detector holds the class index for every pixel, with -1 marking pixels that
belong to no zone:
zones = processor.segmented_detector
import matplotlib.pyplot as plt
plt.imshow(zones.cpu())
plt.title('Detector zones')
plt.colorbar(label='Class')A full classification pipeline
import torch
from svetlanna import SimulationParameters, LinearOpticalSetup
from svetlanna.elements import ThinLens, FreeSpace, DiffractiveLayer
from svetlanna.detector import Detector, DetectorProcessorClf
from svetlanna.transforms import ToWavefront
from svetlanna.units import ureg
params = SimulationParameters.from_ranges(
x_range=(-5*ureg.mm, 5*ureg.mm), x_points=256,
y_range=(-5*ureg.mm, 5*ureg.mm), y_points=256,
wavelength=632.8*ureg.nm
)
setup = LinearOpticalSetup([
DiffractiveLayer(params, mask=torch.rand(256, 256) * 2 * torch.pi),
FreeSpace(params, distance=50*ureg.mm, method='zpASM'),
ThinLens(params, focal_length=100*ureg.mm),
FreeSpace(params, distance=100*ureg.mm, method='zpASM'),
])
detector = Detector(params, func='intensity')
processor = DetectorProcessorClf(
num_classes=10,
simulation_parameters=params,
segmentation_type='strips'
)
def classify(image):
"""Classify a normalised image in [0, 1]."""
wf = ToWavefront(modulation_type='phase')(image) # image → wavefront
wf_out = setup(wf) # optical processing
intensity = detector(wf_out) # detection
return processor(intensity) # class scores
image = torch.rand(256, 256)
scores = classify(image)
print(f"Predicted class: {scores.argmax(dim=1).item()}")Training a classifier
import torch.optim as optim
import torch.nn.functional as F
class OpticalClassifier(torch.nn.Module):
def __init__(self, params, num_classes, grid=256):
super().__init__()
mask = torch.nn.Parameter(torch.rand(grid, grid) * 2 * torch.pi)
self.setup = LinearOpticalSetup([
DiffractiveLayer(params, mask=mask),
FreeSpace(params, distance=50*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, wf):
wf = self.setup(wf)
intensity = self.detector(wf)
return self.processor.batch_forward(intensity)
model = OpticalClassifier(params, num_classes=10)
optimizer = optim.Adam(model.parameters(), lr=0.01)
for epoch in range(100):
optimizer.zero_grad()
logits = model(wf_input)
loss = F.cross_entropy(logits, labels)
loss.backward()
optimizer.step()
if epoch % 10 == 0:
print(f"Epoch {epoch}: loss = {loss.item():.4f}")The zone integrals are not normalised probabilities. Feed them to cross_entropy as
logits, or normalise them yourself before interpreting them as probabilities.
See also
- LinearOpticalSetup — building the optical part
- Transforms — preparing input data
- Optimisation — training loops