Training without a reference

Training without a reference#

Open in Colab

Aim. Train the unrolled network of Staged training of an unrolled network from undersampled k-space alone, with no fully sampled reference, and measure how much of the supervised network’s image quality it retains.

Dynamic and high-dimensional acquisitions – a cine, a functional run, a fingerprinting series – are rarely acquired fully sampled: undersampling is what makes them feasible within a breath-hold or a scan time. There is then no reference image to train against. Self-supervised learning via data undersampling (SSDU) [1] trains against the measured k-space itself. The acquired phase encodes \(\Omega\) are split into two disjoint sets, \(\Theta\) and \(\Lambda\); the network reconstructs from \(\Theta\), and the loss compares the k-space of its reconstruction with the measured data on the held-out set \(\Lambda\),

\[\mathcal{L} = \frac{\|y_\Lambda - A_\Lambda f_\theta(y_\Theta)\|_2}{\|y_\Lambda\|_2} + \frac{\|y_\Lambda - A_\Lambda f_\theta(y_\Theta)\|_1}{\|y_\Lambda\|_1},\]

where \(A_\Lambda\) is the SENSE encoding restricted to \(\Lambda\). A new split is drawn at every step, so over the training every acquired line is both reconstructed from and held out [2]. At inference the network reconstructs from all of \(\Omega\).

Learning objectives

  • Partition the acquired phase encodes with bartorch.learning.split().

  • Train an unrolled network self-supervised with bartorch.learning.Reconstruction, by giving items the sampling pattern instead of a reference.

  • Compare with the same network trained against references, and with CG-SENSE.

It follows Staged training of an unrolled network. The next lesson, Annealed plug-and-play, uses a denoiser trained once for any acquisition.

import csv
import logging
from pathlib import Path

import brainweb_dl
import lightning
import numpy as np
import torch
from brainweb_dl import get_mri
from monai.metrics import PSNRMetric, SSIMMetric
from torch.utils.data import DataLoader

import bartorch
import bartorch.tools as bt
from bartorch import learning, linop, optim, priors

SIZE = 96
COILS = 8
ITERATIONS = 4
ACCELERATION = 4
EPOCHS = 16

_ = torch.manual_seed(0)

Data#

The slices, coils and fourfold undersampling of Staged training of an unrolled network: subject 0 to train on and subject 4 to validate on. The references are kept only to score the results; the self-supervised network never sees them.

train_images = brain_slices(subject=0, count=24)
valid_images = brain_slices(subject=4, count=8)

sensitivities = bt.coils(t=bt.grid(D=(SIZE, SIZE, 1)), n=COILS)[:, 0]
sensitivities = sensitivities / bartorch.rss(sensitivities, axes=(0,), keepdim=True)

density = torch.exp(-0.5 * ((torch.arange(SIZE) - SIZE / 2) / (SIZE / 6)) ** 2)
lines = torch.rand(SIZE, generator=torch.Generator().manual_seed(1)) < density / density.sum() * (
    SIZE / ACCELERATION
)
lines[SIZE // 2 - 4 : SIZE // 2 + 4] = True
pattern = lines.to(torch.complex64)[:, None].expand(SIZE, SIZE).contiguous()
A = linop.CartesianSense(sensitivities, (SIZE, SIZE), pattern=pattern)
NOISE = 0.02

generator = torch.Generator().manual_seed(3)
kspace = {
    "train": [
        A(x) + NOISE * torch.randn(A.oshape, dtype=torch.complex64, generator=generator)
        for x in train_images
    ],
    "valid": [
        A(x) + NOISE * torch.randn(A.oshape, dtype=torch.complex64, generator=generator)
        for x in valid_images
    ],
}

The split#

The readout is fully sampled, so the unit of the split is the phase-encode line: the pattern given to split() has one entry per line and broadcasts over the coils and the readout. A quarter of the acquired lines are held out, drawn with a Gaussian density across \(k_y\), and the eight central lines always stay in \(\Theta\): a reconstruction without the centre of k-space would lose the image contrast, and the loss would be dominated by it.

acquired = lines.float()[:, None]
keep, held = learning.split(acquired, 0.25, keep=(8, 1), generator=torch.Generator().manual_seed(0))
print(
    f"{int(acquired.sum())} acquired lines: {int(keep.sum())} to reconstruct from, "
    f"{int(held.sum())} held out"
)
05 self supervised training
24 acquired lines: 20 to reconstruct from, 4 held out

Two networks, one trained each way#

Both are the iteration-conditioned unrolled network of the previous lesson, trained end to end for the same number of epochs from the same initialization; only the items differ. A self-supervised item carries the acquired pattern in place of a target, and Reconstruction then draws a split at every step. Its validation loss is the held-out loss on a split fixed for the whole run, and needs no reference either.

def unrolled():
    torch.manual_seed(0)
    network = learning.UNet(2, spatial=2, widths=(16, 32, 64), steps=True)
    prior = priors.ImplicitPrior(learning.ComplexNet(network, spatial=2), step=True)
    block = optim.ISTBlock(prior, step=1.0)
    block.step.requires_grad_()
    return learning.Unrolled(block, iterations=ITERATIONS, checkpoint=True)


def items(part, supervised):
    images = train_images if "train" == part else valid_images
    made = []
    for x, y in zip(images, kspace[part]):
        item = {"y": y, "A": A}
        item.update({"target": x} if supervised else {"pattern": acquired})
        made.append(item)
    return made


models = {}
for name, supervised in (("supervised", True), ("self-supervised", False)):
    models[name] = unrolled()
    trainer = lightning.Trainer(
        max_epochs=EPOCHS,
        accelerator="cpu",
        logger=False,
        enable_checkpointing=False,
        enable_model_summary=False,
        enable_progress_bar=False,
    )
    trainer.fit(
        learning.Reconstruction(
            models[name], "end-to-end", lr=1e-3, fraction=0.25, split_options={"keep": (8, 1)}
        ),
        DataLoader(items("train", supervised), batch_size=4, shuffle=True, collate_fn=list),
        DataLoader(items("valid", supervised), batch_size=4, collate_fn=list),
    )
/opt/hostedtoolcache/Python/3.12.14/x64/lib/python3.12/site-packages/lightning/pytorch/utilities/_pytree.py:21: `isinstance(treespec, LeafSpec)` is deprecated, use `isinstance(treespec, TreeSpec) and treespec.is_leaf()` instead.
/opt/hostedtoolcache/Python/3.12.14/x64/lib/python3.12/site-packages/lightning/pytorch/trainer/connectors/data_connector.py:434: The 'val_dataloader' does not have many workers which may be a bottleneck. Consider increasing the value of the `num_workers` argument` to `num_workers=3` in the `DataLoader` to improve performance.
/opt/hostedtoolcache/Python/3.12.14/x64/lib/python3.12/site-packages/lightning/pytorch/trainer/connectors/data_connector.py:434: The 'train_dataloader' does not have many workers which may be a bottleneck. Consider increasing the value of the `num_workers` argument` to `num_workers=3` in the `DataLoader` to improve performance.

Results#

Both networks reconstruct from all acquired lines of the validation subject and are scored against its reference.

psnr = PSNRMetric(max_val=1.0)
ssim = SSIMMetric(spatial_dims=2, data_range=1.0)
truth = torch.stack(valid_images).abs()[:, None]

with torch.no_grad():
    results = {name: torch.stack([m(y, A) for y in kspace["valid"]]) for name, m in models.items()}
results["CG SENSE, 20 iterations"] = torch.stack(
    [optim.cg(y, A, maxiter=20) for y in kspace["valid"]]
)
for name, made in results.items():
    a = made.abs()[:, None]
    print(
        f"{name:>24}   PSNR {float(psnr(a, truth).mean()):5.2f} dB   "
        f"SSIM {float(ssim(a, truth).mean()):.3f}"
    )
             supervised   PSNR 30.20 dB   SSIM 0.945
        self-supervised   PSNR 28.59 dB   SSIM 0.726
CG SENSE, 20 iterations   PSNR 24.27 dB   SSIM 0.558

The self-supervised network is trained on less information: each step reconstructs from three quarters of the acquired lines and is told nothing about the lines never acquired, which is where the supervised network learns most. The gap between the two is what a fully sampled reference would have bought; the self-supervised network needs nothing beyond the data the protocol already acquires.

In the images below, both networks suppress the noise that the CG-SENSE unfolding amplifies across the whole field of view. The self-supervised network keeps more residual aliasing along the phase-encode direction (vertical), which its error map shows as horizontal striping: the lines never acquired are the ones it cannot score against.

  • reference, CG-SENSE, supervised, self-supervised
  • reference, enlarged, CG-SENSE, enlarged, supervised, enlarged, self-supervised, enlarged
  • CG-SENSE NRMSE 0.124, supervised NRMSE 0.063, self-supervised NRMSE 0.076

References#

Total running time of the script: (2 minutes 50.589 seconds)

Gallery generated by Sphinx-Gallery