susan.ml.Noise2NoiseTrainer¶
- class susan.ml.Noise2NoiseTrainer(model: Noise2Noise, loss: str = 'l2')[source]¶
Bases:
objectTraining 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 predicth2fromh1(and vice versa). At inference a single volumevis denoised bypredict(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 → h2andh2 → h1— correct for plain denoising, where both halves share the same signal.With
directional=Trueonly theh1 → h2term is used: slot 0 is treated as the input and slot 1 as the target. Use this for pairs produced byVolumePairs.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.Tensorwith 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.
0disables printing. Default:10.directional (bool) – When
True, train onlyslot 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
Nonethe whole volume is processed at once (same aspredict()). 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