#!/usr/bin/env python3 """Train TinyDigitDiffusion on dynamically composed multi-digit MNIST images.""" from __future__ import annotations import argparse import copy import json import math import random import time from pathlib import Path import torch import torch.nn.functional as F from torch.utils.data import DataLoader, Dataset from torchvision.datasets import MNIST from tqdm import tqdm from tiny_digit_diffusion import ( NULL_TOKEN, PAD_TOKEN, DiffusionSchedule, ModelConfig, TinyDigitDiffusion, count_parameters, ddim_sample, save_prompt_sheet, save_weights, ) class MultiDigitMNIST(Dataset): def __init__(self, root: str | Path, max_digits: int, samples_per_epoch: int, train: bool = True): self.mnist = MNIST(root=str(root), train=train, download=True) self.images = self.mnist.data self.labels = self.mnist.targets self.max_digits = max_digits self.samples_per_epoch = samples_per_epoch def __len__(self) -> int: return self.samples_per_epoch def __getitem__(self, index: int): del index length = int(torch.randint(1, self.max_digits + 1, ()).item()) tokens = torch.full((self.max_digits,), PAD_TOKEN, dtype=torch.long) start_slot = (self.max_digits - length) // 2 canvas = torch.zeros(1, 32, self.max_digits * 32, dtype=torch.float32) chosen = torch.randint(0, len(self.images), (length,)) for offset, image_index in enumerate(chosen): digit = self.images[image_index].float().div(255) label = int(self.labels[image_index]) slot = start_slot + offset tokens[slot] = label x = slot * 32 + 2 + int(torch.randint(-2, 3, ()).item()) y = 2 + int(torch.randint(-2, 3, ()).item()) x = min(max(x, slot * 32), slot * 32 + 4) y = min(max(y, 0), 4) intensity = float(torch.empty(()).uniform_(0.85, 1.0)) canvas[0, y : y + 28, x : x + 28] = torch.maximum( canvas[0, y : y + 28, x : x + 28], digit * intensity ) return canvas.mul(2).sub(1), tokens, torch.tensor(length, dtype=torch.long) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--output-dir", required=True) parser.add_argument("--data-dir", default="data/mnist") parser.add_argument("--epochs", type=int, default=30) parser.add_argument("--samples-per-epoch", type=int, default=60_000) parser.add_argument("--batch-size", type=int, default=64) parser.add_argument("--learning-rate", type=float, default=2e-4) parser.add_argument("--weight-decay", type=float, default=0.0) parser.add_argument("--warmup-steps", type=int, default=500) parser.add_argument("--grad-clip", type=float, default=1.0) parser.add_argument("--condition-dropout", type=float, default=0.1) parser.add_argument("--ema-decay", type=float, default=0.999) parser.add_argument("--num-workers", type=int, default=4) parser.add_argument("--max-steps", type=int, default=0) parser.add_argument("--log-steps", type=int, default=50) parser.add_argument("--sample-every", type=int, default=1) parser.add_argument("--sample-steps", type=int, default=40) parser.add_argument("--guidance-scale", type=float, default=1.0) parser.add_argument("--checkpoint-every", type=int, default=5) parser.add_argument("--seed", type=int, default=1234) parser.add_argument("--device", choices=["auto", "cpu", "cuda"], default="auto") return parser.parse_args() def update_ema(ema_model: torch.nn.Module, model: torch.nn.Module, decay: float) -> None: with torch.no_grad(): for ema_parameter, parameter in zip(ema_model.parameters(), model.parameters()): ema_parameter.lerp_(parameter, 1 - decay) for ema_buffer, buffer in zip(ema_model.buffers(), model.buffers()): ema_buffer.copy_(buffer) def main() -> None: args = parse_args() random.seed(args.seed) torch.manual_seed(args.seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(args.seed) device = torch.device( "cuda" if args.device == "auto" and torch.cuda.is_available() else "cpu" if args.device == "auto" else args.device ) output_dir = Path(args.output_dir).resolve() output_dir.mkdir(parents=True, exist_ok=True) model_dir = output_dir / "model" sample_dir = output_dir / "samples" checkpoint_dir = output_dir / "checkpoints" model_dir.mkdir(exist_ok=True) sample_dir.mkdir(exist_ok=True) checkpoint_dir.mkdir(exist_ok=True) config = ModelConfig() config.save(model_dir / "config.json") model = TinyDigitDiffusion(config).to(device=device, dtype=torch.float32) ema_model = copy.deepcopy(model).requires_grad_(False).eval() parameters = count_parameters(model) print("Device:", device) print("Dtype: float32") print("Parameters:", f"{parameters:,}") print("Image size:", f"{config.image_height}x{config.image_width}") print("Maximum digits:", config.max_digits) dataset = MultiDigitMNIST(args.data_dir, config.max_digits, args.samples_per_epoch) loader = DataLoader( dataset, batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers, pin_memory=device.type == "cuda", persistent_workers=args.num_workers > 0, drop_last=True, ) steps_per_epoch = len(loader) total_steps = args.max_steps if args.max_steps > 0 else args.epochs * steps_per_epoch if total_steps <= args.warmup_steps: args.warmup_steps = max(0, total_steps // 10) print("Steps per epoch:", steps_per_epoch) print("Training steps:", total_steps) optimizer = torch.optim.AdamW(model.parameters(), lr=args.learning_rate, weight_decay=args.weight_decay) def lr_factor(step: int) -> float: if args.warmup_steps and step < args.warmup_steps: return max((step + 1) / args.warmup_steps, 1 / args.warmup_steps) progress = (step - args.warmup_steps) / max(total_steps - args.warmup_steps, 1) return 0.1 + 0.9 * 0.5 * (1 + math.cos(progress * math.pi)) scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_factor) diffusion = DiffusionSchedule(config.diffusion_steps, device) history: list[dict] = [] recent_losses: list[float] = [] global_step = 0 started = time.monotonic() sample_prompts = ["0", "7", "42", "2026", "12345678", "99999999"] model.train() stop = False for epoch in range(1, args.epochs + 1): progress = tqdm(loader, desc=f"epoch {epoch}/{args.epochs}") for clean, tokens, lengths in progress: clean = clean.to(device, non_blocking=True) tokens = tokens.to(device, non_blocking=True) lengths = lengths.to(device, non_blocking=True) drop = torch.rand(len(clean), device=device) < args.condition_dropout tokens = tokens.clone() lengths = lengths.clone() tokens[drop] = NULL_TOKEN lengths[drop] = 0 timesteps = torch.randint(0, config.diffusion_steps, (len(clean),), device=device) noise = torch.randn_like(clean) noisy = diffusion.add_noise(clean, noise, timesteps) predicted = model(noisy, timesteps, tokens, lengths) loss = F.mse_loss(predicted, noise) if not torch.isfinite(loss): raise RuntimeError(f"Non-finite loss at step {global_step + 1}: {loss}") optimizer.zero_grad(set_to_none=True) loss.backward() if args.grad_clip > 0: torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip) optimizer.step() scheduler.step() global_step += 1 effective_ema_decay = min( args.ema_decay, (1 + global_step) / (10 + global_step) ) update_ema(ema_model, model, effective_ema_decay) recent_losses.append(float(loss.detach().cpu())) if global_step % args.log_steps == 0: average = sum(recent_losses[-args.log_steps:]) / min(len(recent_losses), args.log_steps) record = { "step": global_step, "epoch": epoch, "loss": average, "learning_rate": scheduler.get_last_lr()[0], } history.append(record) progress.set_postfix(loss=f"{average:.4f}", lr=f"{record['learning_rate']:.2e}") if global_step >= total_steps: stop = True break if args.sample_every > 0 and (epoch % args.sample_every == 0 or stop): samples = ddim_sample( ema_model, sample_prompts, device, sampling_steps=args.sample_steps, guidance_scale=args.guidance_scale, seed=args.seed + epoch, ) save_prompt_sheet(samples, sample_prompts, sample_dir / f"epoch_{epoch:03d}.png") if args.checkpoint_every > 0 and epoch % args.checkpoint_every == 0 and not stop: torch.save( { "epoch": epoch, "step": global_step, "model": model.state_dict(), "ema_model": ema_model.state_dict(), "optimizer": optimizer.state_dict(), "scheduler": scheduler.state_dict(), "args": vars(args), }, checkpoint_dir / f"epoch_{epoch:03d}.pt", ) (output_dir / "training_history.json").write_text( json.dumps(history, indent=2), encoding="utf-8" ) if stop: break save_weights(ema_model, model_dir / "model.safetensors") elapsed = time.monotonic() - started metadata = { "parameter_count": parameters, "training_steps": global_step, "epochs_completed": epoch, "final_recent_loss": sum(recent_losses[-100:]) / min(len(recent_losses), 100), "training_seconds": elapsed, "args": vars(args), "config": vars(config), "sample_prompts": sample_prompts, "recommended_inference": { "sampling_steps": 50, "guidance_scale": 1.0, }, } (output_dir / "artifact_metadata.json").write_text( json.dumps(metadata, indent=2), encoding="utf-8" ) final_samples = ddim_sample( ema_model, sample_prompts, device, sampling_steps=max(args.sample_steps, 50), guidance_scale=args.guidance_scale, seed=0, ) save_prompt_sheet(final_samples, sample_prompts, output_dir / "final_samples.png") print("Done:", output_dir) print("Final recent loss:", metadata["final_recent_loss"]) print("Training seconds:", elapsed) if __name__ == "__main__": main()