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_file

An exported program (.pt2, from torch.export.save) or a TorchScript file (any other suffix, from torch.jit.save). Export in eval mode; exported modules keep the mode they were exported in. TorchScript is deprecated in recent PyTorch releases, so prefer .pt2 for new models.

model_factory

A Python callable returning an nn.Module, as "package.module:function". With model_factory_file (a .py file, resolved relative to the config like class_file), model_factory is just the function’s name. model_kwargs are passed to it, and state_dict_file (a torch.save of 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:

  1. conversion to float32;

  2. flux_normalization: sum divides by the total flux (the image sums to 1), mean by the mean pixel. A frame whose flux is at most 1e-12 is dark and gives an all-zero signal instead of the model’s response to noise;

  3. sqrt_stretch: negative pixels are clipped to 0 and the square root is taken;

  4. conversion to dtype and reshape to input_shape (default [1, 1, *image_shape]: one single-channel image; set [1, N] for an MLP on the flattened image);

  5. the model; its output is flattened and converted to float32;

  6. output_scale_file: element-wise multiplication by signal_size factors (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_size equals 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 command c applied on a flat wavefront the model should output c (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 wfs stream 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_iters forward 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_active says which path runs;

  • with the common gpu_device key set as well, a GPU-backed wfs stream is attached on the device (no host round trip) and the signal stream is created GPU-backed with a CPU mirror, as SlopesProcess does.

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: Component

Publish a PyTorch model’s output on each WFS image as the loop’s signal.

Sits in the slopes section in place of SlopesProcess: it reads the wfs image stream and writes signal (signal_size float32 values), stamped with each image’s frame id. When signal_2d_shape is set it also writes the output reshaped to that shape to signal_2d for 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 or model_factory.

model_factorystr

"module:function" returning an nn.Module; with model_factory_file, the function’s name in that file.

model_factory_filestr

Python file defining model_factory (resolved like class_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

.npy file of signal_size per-element output factors.

signal_2d_shapelist of int

Shape of an optional signal_2d display 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_device key keeps its usual meaning: it attaches the wfs input on the GPU (when that stream is GPU-backed) and creates GPU-backed output streams. device is 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 wfs read 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 ValueError when the model’s output does not have signal_size elements.

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 of image_shape) and returns the model output as a flat float32 vector of signal_size elements. Every call goes through the same path:

  1. 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.

  2. forward() runs on the device buffer: conversion to float32, optional flux normalisation and square-root stretch, conversion to the model dtype, reshape to input_shape, the model, conversion back to float32 and the optional per-element output scale. With cuda_graph this whole step is one captured CUDA graph.

  3. 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 device and dtype and put in eval mode. It must return one tensor with signal_size elements.

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 signal stream length).

devicestr

"cpu", "cuda" or "cuda:N".

dtype{“float32”, “float16”}

Model precision. float16 needs 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 below FLUX_EPS gives an all-zero output.

sqrt_stretchbool

Clip negative pixels to 0 and take the square root (after normalisation).

output_scalearray_like, optional

signal_size factors 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 (see graph_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_size elements. 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 a model_factory called with model_kwargs that returns an nn.Module. state_dict_file (a torch.save of a state dict, loaded with weights_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)