Training a reconstruction model#

This example provides a very simple quick start introduction to training reconstruction networks with DeepInverse for solving imaging inverse problems.

Training requires these components, all of which you can define with DeepInverse:

Here, we demonstrate a simple experiment of training a UNet on an inpainting task on the Urban100 dataset of natural images.

import deepinv as dinv
import torch

device = dinv.utils.get_device()
rng = torch.Generator(device=device).manual_seed(0)
Selected GPU 0 with 21592.4375 MiB free memory

Setup#

First, define the physics that we want to train on.

physics = dinv.physics.Inpainting((1, 64, 64), mask=0.8, device=device, rng=rng)

Then define the dataset. Here we simulate a dataset of measurements from Urban100.

Tip

See datasets for types of datasets DeepInverse supports: e.g. paired, ground-truth-free, single-image…

from torchvision.transforms import Compose, ToTensor, Resize, CenterCrop, Grayscale

dataset = dinv.datasets.Urban100HR(
    dinv.utils.get_cache_home() / "datasets" / "Urban100",
    download=True,
    transform=Compose([ToTensor(), Grayscale(), Resize(256), CenterCrop(64)]),
)

train_dataset, test_dataset = torch.utils.data.random_split(
    torch.utils.data.Subset(dataset, range(50)), (0.8, 0.2)
)

dataset_path = dinv.datasets.generate_dataset(
    train_dataset=train_dataset,
    test_dataset=test_dataset,
    physics=physics,
    device=device,
    save_dir=".",
    batch_size=1,
)

train_dataloader = torch.utils.data.DataLoader(
    dinv.datasets.HDF5Dataset(dataset_path, train=True), shuffle=True
)
test_dataloader = torch.utils.data.DataLoader(
    dinv.datasets.HDF5Dataset(dataset_path, train=False), shuffle=False
)
Dataset has been saved at ./dinv_dataset0.h5

Visualize a data sample:

x, y = next(iter(test_dataloader))
dinv.utils.plot({"Ground truth": x, "Measurement": y, "Mask": physics.mask})
Ground truth, Measurement, Mask

For the model we use an artifact removal model, where \(\phi_{\theta}\) is a U-Net

\[f_{\theta}(y) = \phi_{\theta}(A^{\top}(y))\]
model = dinv.models.ArtifactRemoval(
    dinv.models.UNet(1, 1, scales=2, batch_norm=False).to(device)
)

Train the model#

We train the model using the deepinv.Trainer class, which cleanly handles all steps for training.

We perform supervised learning and use the mean squared error as loss function. See losses for all supported state-of-the-art loss functions.

We evaluate using the PSNR metric. See metrics for all supported metrics.

Note

In this example, we only train for a few epochs to keep the training time short. For a good reconstruction quality, we recommend to train for at least 100 epochs.

trainer = dinv.Trainer(
    model=model,
    physics=physics,
    optimizer=torch.optim.Adam(model.parameters(), lr=1e-3),
    train_dataloader=train_dataloader,
    eval_dataloader=test_dataloader,
    epochs=5,
    losses=dinv.loss.SupLoss(metric=dinv.metric.MSE()),
    metrics=dinv.metric.PSNR(),
    device=device,
    plot_images=True,
    show_progress_bar=False,
)

_ = trainer.train()
  • Ground truth, Measurement, Reconstruction
  • Ground truth, Measurement, Reconstruction
  • Ground truth, Measurement, Reconstruction
  • Ground truth, Measurement, Reconstruction
  • Ground truth, Measurement, Reconstruction
  • Ground truth, Measurement, Reconstruction
  • Ground truth, Measurement, Reconstruction
  • Ground truth, Measurement, Reconstruction
  • Ground truth, Measurement, Reconstruction
  • Ground truth, Measurement, Reconstruction
/local/jtachell/deepinv/deepinv/deepinv/training/trainer.py:1337: UserWarning: non_blocking_transfers=True but DataLoader.pin_memory=False; set pin_memory=True to overlap host-device copies with compute.
  self.setup_train()
The model has 443585 trainable parameters
Train epoch 0: TotalLoss=0.02, PSNR=18.294
Eval epoch 0: PSNR=22.879
Best model saved at epoch 1
Train epoch 1: TotalLoss=0.004, PSNR=24.917
Eval epoch 1: PSNR=27.499
Best model saved at epoch 2
Train epoch 2: TotalLoss=0.002, PSNR=27.746
Eval epoch 2: PSNR=29.611
Best model saved at epoch 3
Train epoch 3: TotalLoss=0.002, PSNR=28.471
Eval epoch 3: PSNR=28.955
Train epoch 4: TotalLoss=0.001, PSNR=29.877
Eval epoch 4: PSNR=31.335
Best model saved at epoch 5

Test the network#

We can now test the trained network using the deepinv.test() function.

The testing function will compute metrics and plot and save the results.

Ground truth, Measurement, No learning, Reconstruction
/local/jtachell/deepinv/deepinv/deepinv/training/trainer.py:1529: UserWarning: non_blocking_transfers=True but DataLoader.pin_memory=False; set pin_memory=True to overlap host-device copies with compute.
  self.setup_train(train=False)
Eval epoch 0: PSNR=31.335, PSNR no learning=14.133
Test results:
PSNR no learning: 14.133 +- 2.443
PSNR: 31.335 +- 3.267

{'PSNR no learning': 14.132827854156494, 'PSNR no learning_std': 2.4431689189965944, 'PSNR': 31.335296821594238, 'PSNR_std': 3.2674640839779805}

Total running time of the script: (0 minutes 10.410 seconds)

Gallery generated by Sphinx-Gallery