Skip to content

trainer

trainer

Image-scale training utilities: streaming datasets, coupling strategies, and a mixed-precision training loop with checkpoint/resume.

BaseCoupling

Bases: ABC

Given a batch x1 of data samples, return a paired (x0, x1).

IndependentCoupling

Bases: BaseCoupling

x0 ~ N(0, I) drawn independently of x1. Default flow-matching setup.

OTCoupling

Bases: BaseCoupling

Mini-batch optimal-transport coupling on squared-L2 cost.

Draws x0 ~ N(0, I) and then permutes it within the batch so that each (x0, x1) pair minimises the batch's total transport cost. Uses scipy.optimize.linear_sum_assignment if available, otherwise falls back to a deterministic greedy nearest-neighbour matching.

ImageFolderStream

ImageFolderStream(root, image_size=None, mode='RGB', transform=None, extensions=_IMAGE_EXTENSIONS)

Bases: Dataset

Streaming dataset over a directory of images.

Parameters:

Name Type Description Default
root Union[str, Path]

directory containing image files (searched recursively).

required
image_size Optional[int]

if given, images are resized to (image_size, image_size).

None
mode str

PIL mode to convert to ("L" for grayscale, "RGB" for colour, "F" for float).

'RGB'
transform Optional[Callable]

optional PIL-image transform applied before tensor conversion. If it returns a torch.Tensor, that tensor is returned as-is (i.e. no default rescaling is applied).

None
extensions Sequence[str]

file extensions to include.

_IMAGE_EXTENSIONS
Source code in deltaflow/trainer/data.py
def __init__(
    self,
    root: Union[str, Path],
    image_size: Optional[int] = None,
    mode: str = "RGB",
    transform: Optional[Callable] = None,
    extensions: Sequence[str] = _IMAGE_EXTENSIONS,
):
    if Image is None:
        raise ImportError("Pillow is required for ImageFolderStream (pip install pillow)")
    self.root = Path(root)
    self.image_size = image_size
    self.mode = mode
    self.transform = transform
    exts = tuple(e.lower() for e in extensions)
    self.paths = sorted(p for p in self.root.rglob("*") if p.suffix.lower() in exts)
    if not self.paths:
        raise ValueError(f"No images with extensions {exts} found under {self.root}")

TrainConfig dataclass

TrainConfig(max_steps=1000, grad_accum_steps=1, mixed_precision=True, amp_dtype=None, grad_clip=1.0, log_every=50, checkpoint_every=500, checkpoint_dir='checkpoints', ema_beta=0.999, device=None, resume_from=None, log_fn=print)

Configuration for train.

Attributes:

Name Type Description
max_steps int

total optimizer steps to run (after resume_from).

grad_accum_steps int

accumulate gradients over this many minibatches before each optimizer step.

mixed_precision bool

if True, run forward+backward under torch.amp.autocast and scale the loss with torch.amp.GradScaler. On CUDA uses bfloat16 if supported, else float16. On CPU it is a no-op.

amp_dtype Optional[dtype]

override the autocast dtype. If None, chosen from the runtime.

grad_clip Optional[float]

max L2 norm for gradient clipping, None to disable.

log_every int

print a training-metric line every N optimizer steps.

checkpoint_every int

write a checkpoint every N optimizer steps.

checkpoint_dir Union[str, Path]

directory to write checkpoints into. Created if missing.

ema_beta Optional[float]

if not None, maintain an EMA copy of the model with this decay rate.

device Optional[str]

torch device string, defaults to CUDA if available.

resume_from Optional[Union[str, Path]]

optional path to a checkpoint saved by this loop.

build_loader

build_loader(dataset, batch_size, num_workers=2, shuffle=True, pin_memory=True, drop_last=True)

Convenience wrapper around torch.utils.data.DataLoader with sensible defaults for image-scale flow-matching training.

Source code in deltaflow/trainer/data.py
def build_loader(
    dataset: Dataset,
    batch_size: int,
    num_workers: int = 2,
    shuffle: bool = True,
    pin_memory: bool = True,
    drop_last: bool = True,
) -> DataLoader:
    """Convenience wrapper around `torch.utils.data.DataLoader` with
    sensible defaults for image-scale flow-matching training."""
    return DataLoader(
        dataset,
        batch_size=batch_size,
        num_workers=num_workers,
        shuffle=shuffle,
        pin_memory=pin_memory,
        drop_last=drop_last,
        persistent_workers=num_workers > 0,
    )

load_checkpoint

load_checkpoint(path, model, optimizer=None, ema_model=None, scaler=None, map_location=None)

Load a checkpoint written by save_checkpoint. Returns the step.

Source code in deltaflow/trainer/loop.py
def load_checkpoint(
    path: Union[str, Path],
    model: nn.Module,
    optimizer: Optional[Optimizer] = None,
    ema_model: Optional[nn.Module] = None,
    scaler: Optional[torch.amp.GradScaler] = None,
    map_location: Optional[Union[str, torch.device]] = None,
) -> int:
    """Load a checkpoint written by `save_checkpoint`. Returns the step."""
    ckpt = torch.load(path, map_location=map_location, weights_only=False)
    model.load_state_dict(ckpt["model"])
    if optimizer is not None and "optimizer" in ckpt:
        optimizer.load_state_dict(ckpt["optimizer"])
    if ema_model is not None and "ema_model" in ckpt:
        ema_model.load_state_dict(ckpt["ema_model"])
    if scaler is not None and "scaler" in ckpt:
        scaler.load_state_dict(ckpt["scaler"])
    return int(ckpt.get("step", 0))

save_checkpoint

save_checkpoint(path, model, optimizer, step, ema_model=None, scaler=None)

Write a checkpoint dictionary to path (creates parent dir).

Source code in deltaflow/trainer/loop.py
def save_checkpoint(
    path: Union[str, Path],
    model: nn.Module,
    optimizer: Optimizer,
    step: int,
    ema_model: Optional[nn.Module] = None,
    scaler: Optional[torch.amp.GradScaler] = None,
) -> None:
    """Write a checkpoint dictionary to ``path`` (creates parent dir)."""
    path = Path(path)
    path.parent.mkdir(parents=True, exist_ok=True)
    ckpt = {
        "model": model.state_dict(),
        "optimizer": optimizer.state_dict(),
        "step": step,
    }
    if ema_model is not None:
        ckpt["ema_model"] = ema_model.state_dict()
    if scaler is not None:
        ckpt["scaler"] = scaler.state_dict()
    torch.save(ckpt, path)

train

train(model, optimizer, loss_fn, dataloader, config=None)

Train model in-place under loss_fn for config.max_steps steps.

Returns the (possibly EMA-shadowed) model that was trained. The original model is always updated in place. If EMA is enabled, an EMA copy is also kept and written into checkpoints, and can be retrieved from the latest checkpoint's "ema_model" key.

Source code in deltaflow/trainer/loop.py
def train(
    model: nn.Module,
    optimizer: Optimizer,
    loss_fn: BaseLoss,
    dataloader: DataLoader,
    config: Optional[TrainConfig] = None,
) -> nn.Module:
    """Train ``model`` in-place under ``loss_fn`` for ``config.max_steps`` steps.

    Returns the (possibly EMA-shadowed) model that was trained. The original
    ``model`` is always updated in place. If EMA is enabled, an EMA copy is
    also kept and written into checkpoints, and can be retrieved from the
    latest checkpoint's ``"ema_model"`` key.
    """
    cfg = config or TrainConfig()
    device = _resolve_device(cfg.device)
    model.to(device)

    amp_dtype = _select_amp_dtype(device, cfg.amp_dtype)
    use_amp = cfg.mixed_precision and device.type == "cuda"
    scaler = torch.amp.GradScaler("cuda", enabled=use_amp and amp_dtype == torch.float16)

    ema = EMA(beta=cfg.ema_beta) if cfg.ema_beta is not None else None
    ema_model: Optional[nn.Module] = None
    if ema is not None:
        ema_model = copy.deepcopy(model).eval().requires_grad_(False)

    start_step = 0
    if cfg.resume_from is not None:
        start_step = load_checkpoint(
            cfg.resume_from,
            model=model,
            optimizer=optimizer,
            ema_model=ema_model,
            scaler=scaler if use_amp else None,
            map_location=device,
        )
        cfg.log_fn(f"[deltaflow.trainer] resumed at step {start_step} from {cfg.resume_from}")

    ckpt_dir = Path(cfg.checkpoint_dir)
    data_iter = _infinite(dataloader)

    model.train()
    step = start_step
    accum_i = 0
    running_loss = 0.0
    t_last = time.time()

    optimizer.zero_grad(set_to_none=True)

    while step < cfg.max_steps:
        batch = next(data_iter)
        x1 = _extract_x1(batch).to(device, non_blocking=True)

        with torch.amp.autocast(
            device_type=device.type,
            dtype=amp_dtype,
            enabled=use_amp,
        ):
            loss = loss_fn(model, x1)
            loss = loss / cfg.grad_accum_steps

        if use_amp and amp_dtype == torch.float16:
            scaler.scale(loss).backward()
        else:
            loss.backward()

        running_loss += loss.item() * cfg.grad_accum_steps
        accum_i += 1

        if accum_i < cfg.grad_accum_steps:
            continue
        accum_i = 0

        if cfg.grad_clip is not None:
            if use_amp and amp_dtype == torch.float16:
                scaler.unscale_(optimizer)
            torch.nn.utils.clip_grad_norm_(model.parameters(), cfg.grad_clip)

        if use_amp and amp_dtype == torch.float16:
            scaler.step(optimizer)
            scaler.update()
        else:
            optimizer.step()
        optimizer.zero_grad(set_to_none=True)

        if ema_model is not None:
            ema.update_model_average(ema_model, model)

        step += 1

        if step % cfg.log_every == 0:
            avg_loss = running_loss / (cfg.log_every * cfg.grad_accum_steps)
            dt = time.time() - t_last
            cfg.log_fn(
                f"[deltaflow.trainer] step {step:>7d} | loss {avg_loss:.4f} | {dt:.1f}s / {cfg.log_every} steps"
            )
            running_loss = 0.0
            t_last = time.time()

        if step % cfg.checkpoint_every == 0:
            save_checkpoint(
                ckpt_dir / f"step_{step:07d}.pt",
                model=model,
                optimizer=optimizer,
                step=step,
                ema_model=ema_model,
                scaler=scaler if use_amp else None,
            )

    save_checkpoint(
        ckpt_dir / "last.pt",
        model=model,
        optimizer=optimizer,
        step=step,
        ema_model=ema_model,
        scaler=scaler if use_amp else None,
    )
    return ema_model if ema_model is not None else model