Download train_tiny_digit_diffusion.py from shibatch/tinydigitdiffusion3m: direct link, hf CLI and curl.
- Browser
- Download file 11.1 kB
-
https://huggingface.co/shibatch/tinydigitdiffusion3m/resolve/main/train_tiny_digit_diffusion.py
- Command line
-
hf download hf://shibatch/tinydigitdiffusion3m/train_tiny_digit_diffusion.py
-
curl -L -o train_tiny_digit_diffusion.py https://huggingface.co/shibatch/tinydigitdiffusion3m/resolve/main/train_tiny_digit_diffusion.py
11.1 kB
| #!/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() | |