tinydigitdiffusion3m / train_tiny_digit_diffusion.py
shibatch's picture
Upload folder using huggingface_hub
58f76dc verified
Raw
History Blame Contribute Delete
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()