Torch Image Reconstructor
TorchImageReconstructor turns wavefront-sensor images into the loop’s
signal with a PyTorch model. Use it for image-based reconstructors, such as a
neural network that maps Shack-Hartmann or pyramid pixels to modes, or a
focal-plane wavefront sensor whose “WFS” is a science-camera image.
It goes in the slopes section in place of SlopesProcess. It reads the
wfs image stream and writes the model output to signal, stamped with
each image’s frame id, so the loop, its calibration methods, telemetry and
manager.latency() work unchanged.
Configuration
slopes:
class_name: TorchImageReconstructor
signal_size: 120 # number of model outputs
model_file: calib/reconstructor.pt2
device: cuda:0 # cpu (default), cuda, cuda:N
dtype: float32 # or float16 (CUDA only)
flux_normalization: sum # none (default), sum or mean
sqrt_stretch: false
output_scale_file: "" # optional .npy, signal_size factors
cuda_graph: true
input_streams: {wfs: wfs}
output_streams: {signal: signal}
functions: [compute_signal]
signal_size is required. At startup the model runs once on a blank image
and the reconstructor raises ValueError if its output does not have
signal_size elements, or if it fails on the input shape. Every key is
listed in the parameter reference below.
Model sources
Set exactly one of:
model_fileAn exported program (
.pt2, fromtorch.export.save) or a TorchScript file (any other suffix, fromtorch.jit.save). Export in eval mode; exported modules keep the mode they were exported in. TorchScript is deprecated in recent PyTorch releases, so prefer.pt2for new models.model_factoryA Python callable returning an
nn.Module, as"package.module:function". Withmodel_factory_file(a.pyfile, resolved relative to the config likeclass_file),model_factoryis just the function’s name.model_kwargsare passed to it, andstate_dict_file(atorch.saveof a state dict) is then loaded strictly.
slopes:
class_name: TorchImageReconstructor
signal_size: 120
model_factory: build_cnn
model_factory_file: models.py
model_kwargs: {num_outputs: 120}
state_dict_file: calib/cnn_weights.pt
functions: [compute_signal]
In a soft-RTC session a model can also be passed directly, or swapped on a running system (the new model is validated, warmed up and captured before it replaces the old one between frames):
reconstructor = TorchImageReconstructor(conf, model=my_module)
reconstructor.set_model(retrained_module)
The model is moved to device and dtype in place.
Preprocessing
The WFS component has already subtracted the dark (and background). On the model’s device, each frame then goes through, in order:
conversion to float32;
flux_normalization:sumdivides by the total flux (the image sums to 1),meanby the mean pixel. A frame whose flux is at most1e-12is dark and gives an all-zero signal instead of the model’s response to noise;sqrt_stretch: negative pixels are clipped to 0 and the square root is taken;conversion to
dtypeand reshape toinput_shape(default[1, 1, *image_shape]: one single-channel image; set[1, N]for an MLP on the flattened image);the model; its output is flattened and converted to float32;
output_scale_file: element-wise multiplication bysignal_sizefactors (for example the inverse of per-mode training normalisation).
The image is the stream array as stored, (height, width) = [y, x];
train on frames read from the wfs stream to match. A model trained on
pyrtc 1.x frames from an adapter that transposed them (GenICam,
Micro-Manager) sees transposed images now; retrain it, or wrap it in a
module that transposes the image’s last two axes first (see
Migrating to pyrtc 2.0).
Plugging into the loop
The loop sizes its interaction and control matrices from the signal
stream, so nothing changes in the loop section. Two common set-ups:
The model outputs modal coefficients (
signal_sizeequals the number of controlled modes, in the corrector’s modal basis and units). Load an identity interaction matrix, so the control matrix is the identity too:from pyrtc.calibration import save_calibration save_calibration("calib/identity_im.npz", np.eye(num_modes, dtype=np.float32), "interaction_matrix")
loop: im_file: calib/identity_im.npz gain: 0.3 functions: [standard_integrator]
The sign convention is the loop’s: it subtracts
gain * CM @ signal, so for a corrector commandcapplied on a flat wavefront the model should outputc(the response an identity interaction matrix describes).The model outputs any other signal, such as a feature vector. Calibrate the interaction matrix through the live pipeline (
loop.compute_im()) exactly as with slopes.
No signal_2d stream is created unless signal_2d_shape is set (it must
hold signal_size elements).
Real-time path and GPU use
TorchModelRunner holds the model and the per-frame path. On a CUDA
device:
frames are read from the
wfsstream straight into a pinned host buffer, copied to a static device buffer on the runner’s own CUDA stream, and the output is copied back into a pinned buffer before one synchronisation;after
warmup_itersforward passes, preprocessing, the model and the output scaling are captured in one CUDA graph (cuda_graph: true). The graph is checked against the eager result on a random frame. A model that cannot be captured (one that synchronises with the host, e.g. with.item(), or has data-dependent control flow) falls back to eager execution with a warning;reconstructor.runner.graph_activesays which path runs;with the common
gpu_devicekey set as well, a GPU-backedwfsstream is attached on the device (no host round trip) and thesignalstream is created GPU-backed with a CPU mirror, asSlopesProcessdoes.
On the CPU the model runs eagerly with torch.inference_mode.
cpu_threads sets torch.set_num_threads, which is process-wide. Hard-RTC
children start with OMP_NUM_THREADS=1 (unless you set it), so a CPU model
there runs on one thread unless cpu_threads asks for more.
Timing
last_compute_time is the time in seconds the last frame took from the end
of the wfs read to the output being on the host (preprocessing, transfers,
model). timing_stats() returns the count, mean, median, p99 and max over the
last timing_window frames, and reset_timing() clears them. In hard-RTC
mode, call them through the launcher (launcher.run("timing_stats")).
Stream handoffs come on top; manager.latency() measures the whole
WFS-to-DM path.
benchmarks/image_reconstructor_bench.py times the runner for a 64x64
image and 120 outputs with a ~15M-parameter CNN and a ~1.1M-parameter MLP.
Results vary by host and GPU load; on an RTX 4060 with a Neoverse-N1 host
(PyTorch 2.11, three runs, microseconds):
Model |
Path |
Median |
p99 |
|---|---|---|---|
CNN, 14.8M |
CPU, float32, 16 threads |
4720-5080 |
5250-7790 |
CNN, 14.8M |
eager, float32 |
1210-1440 |
1780-2720 |
CNN, 14.8M |
CUDA graph, float32 |
580-750 |
1150-2700 |
CNN, 14.8M |
CUDA graph, float16 |
380-560 |
1100-2550 |
MLP, 1.1M |
eager, float32 |
610-640 |
920-2260 |
MLP, 1.1M |
CUDA graph, float32 |
90-180 |
360-510 |
Eager execution is bound by kernel-launch overhead for models this size, so the graph roughly halves the CNN’s latency and cuts the MLP’s by 3-6x.
Parameters
- class pyrtc.image_reconstructor.TorchImageReconstructor(conf, model=None)[source]
Bases:
ComponentPublish a PyTorch model’s output on each WFS image as the loop’s signal.
Sits in the
slopessection in place ofSlopesProcess: it reads thewfsimage stream and writessignal(signal_sizefloat32 values), stamped with each image’s frame id. Whensignal_2d_shapeis set it also writes the output reshaped to that shape tosignal_2dfor viewers.Config
- signal_sizeint
Number of model outputs. The model is checked against it at startup.
- model_filestr
Exported program (
.pt2,torch.export.save) or TorchScript file (torch.jit.save, any other suffix). Either this ormodel_factory.- model_factorystr
"module:function"returning annn.Module; withmodel_factory_file, the function’s name in that file.- model_factory_filestr
Python file defining
model_factory(resolved likeclass_file).- model_kwargsdict
Keyword arguments for
model_factory.- state_dict_filestr
State dict loaded (strictly) into the model.
- devicestr
Model device:
"cpu"(default),"cuda"or"cuda:N".- dtypestr
"float32"(default) or"float16"(CUDA only).- input_shapelist of int
Shape the model takes; default
[1, 1, *image_shape].- flux_normalizationstr
"none"(default),"sum"or"mean".- sqrt_stretchbool
Square-root stretch (negative pixels clipped to 0). Default False.
- output_scale_filestr
.npyfile ofsignal_sizeper-element output factors.- signal_2d_shapelist of int
Shape of an optional
signal_2ddisplay stream.- cuda_graphbool
Capture the model in a CUDA graph (CUDA only). Default True.
- warmup_itersint
Forward passes before capture. Default 10.
- cpu_threadsint
When set,
torch.set_num_threads(process-wide) for CPU models.- timing_windowint
Number of recent per-frame compute times kept for
timing_stats(). Default 1000.
The common
gpu_devicekey keeps its usual meaning: it attaches thewfsinput on the GPU (when that stream is GPU-backed) and creates GPU-backed output streams.deviceis where the model runs.Attributes
- runnerTorchModelRunner
The model and its real-time path.
- last_compute_timefloat or None
Seconds spent on the last frame, from the end of the
wfsread to the output being on the host (preprocessing, transfers, model).- frames_processedint
Frames published since construction.
- compute_signal()[source]
Read one WFS image, run the model and publish
signal.
- load_model()[source]
Build the model from the config’s
model_file/model_factory.
- read(block=True)[source]
Read the current signal.
- read_image(block=True)[source]
Read the current WFS image.
- reset_timing()[source]
Forget the recorded compute times.
- Return type:
None
- set_model(model)[source]
Install
model: validate its output size, warm it up and capture it.Safe on a running instance: the new runner is built first and swapped in between frames. Raises
ValueErrorwhen the model’s output does not havesignal_sizeelements.- Return type:
None
- timing_stats()[source]
Return per-frame compute-time statistics (seconds) over the recent window.
- Return type:
dict[str, float]
- class pyrtc.image_reconstructor.TorchModelRunner(model, *, image_shape, image_dtype=<class 'numpy.float32'>, signal_size, device='cpu', dtype='float32', input_shape=None, flux_normalization='none', sqrt_stretch=False, output_scale=None, cuda_graph=True, warmup_iters=10)[source]
Run a PyTorch model on single WFS images, with optional CUDA graphs.
The runner has no streams:
run()takes one image (NumPy array or torch tensor ofimage_shape) and returns the model output as a flat float32 vector ofsignal_sizeelements. Every call goes through the same path:On CUDA, a NumPy image is copied into a pinned host buffer (or written there directly by the caller through
host_input), then copied to a static device buffer on the runner’s own CUDA stream.forward()runs on the device buffer: conversion to float32, optional flux normalisation and square-root stretch, conversion to the model dtype, reshape toinput_shape, the model, conversion back to float32 and the optional per-element output scale. Withcuda_graphthis whole step is one captured CUDA graph.The result is copied into a pinned host buffer and the stream is synchronised.
Parameters
- modeltorch.nn.Module or torch.jit.ScriptModule
The model. It is moved to
deviceanddtypeand put in eval mode. It must return one tensor withsignal_sizeelements.- image_shapetuple of int
Shape of the WFS image as stored in the stream.
- image_dtypenumpy dtype
Dtype of the WFS image.
- signal_sizeint
Number of model outputs (the
signalstream length).- devicestr
"cpu","cuda"or"cuda:N".- dtype{“float32”, “float16”}
Model precision.
float16needs a CUDA device. Outputs are always float32.- input_shapetuple of int, optional
Shape the model takes. Default
(1, 1, *image_shape): one image, one channel. Must hold as many elements as the image.- flux_normalization{“none”, “sum”, “mean”}
Divide the image by its total (
sum) or mean pixel (mean) flux. A frame whose flux is at or belowFLUX_EPSgives an all-zero output.- sqrt_stretchbool
Clip negative pixels to 0 and take the square root (after normalisation).
- output_scalearray_like, optional
signal_sizefactors multiplied into the output.- cuda_graphbool
Capture
forward()in a CUDA graph on CUDA devices. If capture fails, or the graph’s output differs from the eager output, the runner logs a warning and stays eager (seegraph_active).- warmup_itersint
Forward passes run before capture (and before timing).
- forward(raw)[source]
Preprocess
raw(an image tensor on the device) and run the model.Returns a flat float32 tensor of
signal_sizeelements. Everything here stays on the device without host synchronisation, so it can be captured in a CUDA graph.
- host_input
NumPy view of the (pinned) host input buffer. Reading a stream straight into it (
read_stream(..., out=runner.host_input)) saves a copy;run()notices and skips its own.
- run(image)[source]
Process one image and return the output as a float32 NumPy vector.
The returned array is the runner’s (pinned) host output buffer and is overwritten by the next call; copy it to keep it.
- Return type:
ndarray
- run_device(image)[source]
Process one image and return the output tensor on the model device.
The returned tensor is the runner’s static output buffer when a CUDA graph is active: it is overwritten by the next call. On CUDA the work is queued on the runner’s stream and not synchronised; use
run()for a host result.
- synchronize()[source]
Wait for the work queued on the runner’s CUDA stream (no-op on CPU).
- Return type:
None
- Parameters:
image_shape (tuple[int, ...])
signal_size (int)
device (str)
dtype (str)
input_shape (tuple[int, ...] | None)
flux_normalization (str)
sqrt_stretch (bool)
cuda_graph (bool)
warmup_iters (int)
- pyrtc.image_reconstructor.load_torch_model(*, model_file='', model_factory='', model_factory_file='', model_kwargs=None, state_dict_file='')[source]
Build the model a reconstructor config describes (on the CPU).
Exactly one source must be given: a
model_file, either an exported program (.pt2,torch.export.save; export it in eval mode) or a TorchScript file (any other suffix,torch.jit.load), or amodel_factorycalled withmodel_kwargsthat returns annn.Module.state_dict_file(atorch.saveof a state dict, loaded withweights_only=True) is then loaded into it strictly.- Parameters:
model_file (str)
model_factory (str)
model_factory_file (str)
model_kwargs (Mapping[str, Any] | None)
state_dict_file (str)