Skip to content

Diffusion (eval is not a forward pass)

A DDPM over eight Gaussians on a ring. CPU, no downloads, ~40 s.

gatle-ignite train --config=examples/diffusion/configs/diffusion_v0.py
python examples/diffusion/scripts/check_sampling.py     # the proof

Two things here are not standard supervised learning, and both land on existing hooks.

Nothing beyond the walk: dataset · model · optimizer · checkpoints · logging · running

step what this example changes
prep_batch overrides eval_step()
loss dict_of_loss_params: eps (mse)
metrics val_metrics: mode_dist (ring_stats), mode_tv (ring_stats), radius_error (ring_stats)
score_name = valid/mode_dist
score_factor = -1

The target is sampled per batch, not supplied by the dataset

prep_batch draws a timestep and a noise vector and puts them in targets. Nothing in the framework assumes targets come from the dataloader, so a plain MSE term wires up to ("targets", "noise") like any other label.

examples/diffusion/trainer/diffusion_trainer.py
import torch

from examples.diffusion.models.model_utils.diffusion import Schedule
from gatle_ignite import BaseTrainer, to_device


class Trainer(BaseTrainer):
    def __init__(self, local_rank, cfg):
        super().__init__(local_rank, cfg)
        # Tiny tensors: 8 threads measured 0.30 s/step against 0.05 at 1. Set here, not in the
        # config, which is imported in the launcher's process rather than this worker's.
        threads = cfg.diffusion.get("torch_threads")
        if threads:
            torch.set_num_threads(threads)
        self.schedule = Schedule(num_steps=cfg.diffusion["num_steps"])

    def prep_batch(self, batch, split="train", **kwargs):
        (x0,) = to_device(batch)
        self.schedule.to(self.device)

        # Drawn per batch, so the same x_0 meets a new noise level every epoch.
        t = self.schedule.sample_timesteps(x0.shape[0], self.device)
        noise = torch.randn_like(x0)
        x_t = self.schedule.q_sample(x0, t, noise)

        return {
            "model_input": {"x_t": x_t, "t": t},
            # The MSE regresses onto `noise`; `x0` is the eval metric's reference.
            "targets": {"noise": noise, "x0": x0},
        }

    def eval_step(self, engine, batch, split="valid"):
        """Generate, don't predict: the batch supplies only the reference x_0 and a sample count."""
        engine.state.batch = None
        engine.state.output = None
        self.model.eval()

        x = self.prep_batch(batch, split=split)
        real = x["targets"]["x0"]

        with torch.no_grad():
            with torch.autocast(
                device_type=self.device_type, dtype=self.dtype, enabled=self.autocast_enabled
            ):

                def eps_fn(x_t, t):
                    return self.forward({"x_t": x_t, "t": t})["noise_pred"]

                # self.device: sampling starts from noise, with no input to take a device from.
                samples = self.schedule.p_sample_loop(eps_fn, tuple(real.shape), self.device)

        return {"y_pred": {"samples": samples.float()}, "target": x}

Eval runs the reverse process

eval_step is overridden to sample (200 ancestral steps from pure noise) instead of doing a forward pass. This is the normal case for generative models, retrieval and beam search, not an exotic one.

Why the loss is not the evidence

train/loss_eps_avg sits at ~0.31 for the whole run and barely moves. A noise-predictor's loss falls even when sampling is broken, so it tells you nothing. The metrics compare sampled points against the true ring.

check_sampling.py scores four columns, and the two controls are the point:

metric                  real vs real  prior N(0,I)     UNTRAINED       TRAINED
valid/mode_dist               0.0615        0.6084    27524.3694        0.1940
  • UNTRAINED diverges to ~2.7e4, proving the sampler is really driven by the weights.
  • prior N(0,I), a sampler that never denoises, is the sharper control, and it is why the score is valid/mode_dist. On valid/mode_tv pure noise scores 0.055 against real data's own 0.057 floor: isotropic noise spreads across eight modes as evenly as the ring does, so mode_tv cannot distinguish a working sampler from no sampler. mode_dist separates them, 0.19 vs 0.61.

Honest limit: trained mode_dist 0.194 against a 0.062 floor. The samples sit on the right ring at the right radius with the right mode occupancy, and are ~3x fuzzier than real data. It learned the distribution; it did not nail it.