| |
| """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() |
|
|