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_distscore_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.
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:
- 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. Onvalid/mode_tvpure noise scores 0.055 against real data's own 0.057 floor: isotropic noise spreads across eight modes as evenly as the ring does, somode_tvcannot distinguish a working sampler from no sampler.mode_distseparates 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.