susan.ml.Noise2NoiseTrainer

class susan.ml.Noise2NoiseTrainer(model: Noise2Noise, loss: str = 'l2')[source]

Bases: object

Training and inference wrapper for Noise2Noise.

Warning

Experimental. This class is part of the machine-learning subpackage and may change or be removed in a future release.

Implements Noise2Noise training: given a half-map pair (h1, h2) where both share the same underlying signal but have independent noise, the model learns to predict h2 from h1 (and vice versa). At inference a single volume v is denoised by predict(v).

Parameters:
  • model (Noise2Noise) – Network to train.

  • loss (str) – 'l2' (MSE) or 'l1' (MAE). Default: 'l2'.

Training

train(vol_pairs, n_epochs: int = 10, lr: float = 0.0001, batch_size: int = 1, print_every: int = 10, directional: bool = False)[source]

Train the network using Noise2Noise on half-map pairs.

By default each sample (h1, h2) from vol_pairs contributes two symmetric terms, h1 → h2 and h2 → h1 — correct for plain denoising, where both halves share the same signal.

With directional=True only the h1 → h2 term is used: slot 0 is treated as the input and slot 1 as the target. Use this for pairs produced by VolumePairs.populate() with an angular perturbation, where slot 0 is a deliberately misaligned reconstruction and slot 1 is cleanly aligned; the reverse term would train the network to introduce misalignment and must be disabled.

Parameters:
  • vol_pairs (VolumePairs or iterable of (Tensor, Tensor)) – Source of half-map pairs. Each element must be a pair of torch.Tensor with shape (1, D, H, W) or (D, H, W).

  • n_epochs (int) – Number of passes over vol_pairs. Default: 10.

  • lr (float) – Learning rate. Default: 1e-4.

  • batch_size (int) – Number of pairs per gradient step. Default: 1.

  • print_every (int) – Print mean loss every this many epochs. 0 disables printing. Default: 10.

  • directional (bool) – When True, train only slot 0 → slot 1 (input → target) without the symmetric reverse term. Default: False.

Returns:

Mean loss per epoch.

Return type:

list of float

Inference

predict(v: ndarray) → ndarray[source]

Denoise a single volume.

Parameters:

v (numpy.ndarray) – Input volume, shape (D, H, W).

Returns:

Denoised volume, same shape, float32.

Return type:

numpy.ndarray

denoise(v: ndarray, patch_size: int = None, overlap: int = 8) → ndarray[source]

Denoise a volume, optionally with overlapping patches.

When patch_size is None the whole volume is processed at once (same as predict()). Patch mode is useful when GPU memory is limited; overlapping regions are averaged.

Parameters:
  • v (numpy.ndarray) – Input volume, shape (D, H, W).

  • patch_size (int or None) – Cubic patch side length. None = whole volume. Default: None.

  • overlap (int) – Overlap in voxels between adjacent patches. Default: 8.

Returns:

Denoised volume, same shape, float32.

Return type:

numpy.ndarray

mace_step(v: ndarray, alpha: float = 0.5) → ndarray[source]

MACE consensus step: blend input with network prediction.

Returns alpha * predict(v) + (1 - alpha) * v.

Parameters:
  • v (numpy.ndarray) – Input volume.

  • alpha (float) – Blending weight (0 = no denoising, 1 = full prediction). Default: 0.5.

Returns:

Blended volume, float32.

Return type:

numpy.ndarray

Persistence

save(path: str)[source]

Save model weights and trainer state.

Parameters:

path (str) – Destination file path.

load_state(path: str)[source]

Restore weights and trainer state saved with save().

Parameters:

path (str) – Path to a file written by save().