PixelModel v4: tiny latent diffusion (DiT + rectified flow), FID 39.54 / CLIP 28.04 (part 2)
6c9c825 verified | from __future__ import annotations | |
| import argparse, copy, math, os, time | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| from dit import DiT | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--work", default="/root/pm4") | |
| ap.add_argument("--steps", type=int, default=80000) | |
| ap.add_argument("--batch", type=int, default=256) | |
| ap.add_argument("--lr", type=float, default=2e-4) | |
| ap.add_argument("--warmup", type=int, default=1000) | |
| ap.add_argument("--dim", type=int, default=384) | |
| ap.add_argument("--depth", type=int, default=12) | |
| ap.add_argument("--heads", type=int, default=6) | |
| ap.add_argument("--cfg-dropout", type=float, default=0.1) | |
| ap.add_argument("--ema", type=float, default=0.9999) | |
| ap.add_argument("--ckpt-every", type=int, default=5000) | |
| ap.add_argument("--log-every", type=int, default=100) | |
| ap.add_argument("--out", default="/root/pm4/ckpt") | |
| ap.add_argument("--resume", default="") | |
| args = ap.parse_args() | |
| dev = "cuda" | |
| os.makedirs(args.out, exist_ok=True) | |
| lat = torch.from_numpy(np.load(os.path.join(args.work, "latents.npy"))).pin_memory() | |
| seq = torch.from_numpy(np.load(os.path.join(args.work, "text_seq.npy"))).pin_memory() | |
| pool = torch.from_numpy(np.load(os.path.join(args.work, "text_pool.npy"))).pin_memory() | |
| null_seq = torch.from_numpy(np.load(os.path.join(args.work, "null_seq.npy"))).float().to(dev) | |
| null_pool = torch.from_numpy(np.load(os.path.join(args.work, "null_pool.npy"))).float().to(dev) | |
| N = lat.shape[0] | |
| print(f"[train] N={N} latents{lat.shape} text{seq.shape} dim={args.dim} depth={args.depth}", flush=True) | |
| model = DiT(dim=args.dim, depth=args.depth, heads=args.heads).to(dev) | |
| print(f"[train] DiT params = {model.num_params():,}", flush=True) | |
| ema = copy.deepcopy(model).eval() | |
| for p in ema.parameters(): | |
| p.requires_grad_(False) | |
| opt = torch.optim.AdamW(model.parameters(), lr=args.lr, betas=(0.9, 0.99), weight_decay=0.0) | |
| start = 0 | |
| if args.resume and os.path.exists(args.resume): | |
| ck = torch.load(args.resume, map_location=dev) | |
| model.load_state_dict(ck["model"]); ema.load_state_dict(ck["ema"]) | |
| opt.load_state_dict(ck["opt"]); start = ck["step"] | |
| print(f"[train] resumed from step {start}", flush=True) | |
| def lr_at(step): | |
| if step < args.warmup: | |
| return args.lr * step / args.warmup | |
| p = (step - args.warmup) / max(1, args.steps - args.warmup) | |
| return args.lr * (0.1 + 0.9 * 0.5 * (1 + math.cos(math.pi * min(1.0, p)))) | |
| def save(step, tag): | |
| path = os.path.join(args.out, f"{tag}.pt") | |
| torch.save({"model": model.state_dict(), "ema": ema.state_dict(), | |
| "opt": opt.state_dict(), "step": step, | |
| "cfg": {"dim": args.dim, "depth": args.depth, "heads": args.heads}}, path) | |
| print(f"[train] saved {path} @ step {step}", flush=True) | |
| model.train() | |
| t0 = time.time() | |
| run_loss = 0.0 | |
| for step in range(start, args.steps): | |
| for g in opt.param_groups: | |
| g["lr"] = lr_at(step) | |
| idx = torch.randint(0, N, (args.batch,)) | |
| x1 = lat[idx].to(dev, non_blocking=True).float() | |
| ts = seq[idx].to(dev, non_blocking=True).float() | |
| tp = pool[idx].to(dev, non_blocking=True).float() | |
| fm = torch.rand(args.batch, device=dev) < 0.5 | |
| if fm.any(): | |
| x1[fm] = torch.flip(x1[fm], dims=[3]) | |
| drop = torch.rand(args.batch, device=dev) < args.cfg_dropout | |
| if drop.any(): | |
| ts[drop] = null_seq | |
| tp[drop] = null_pool | |
| x0 = torch.randn_like(x1) | |
| u = torch.randn(args.batch, device=dev) | |
| t = torch.sigmoid(u) | |
| tb = t.view(-1, 1, 1, 1) | |
| xt = (1 - tb) * x0 + tb * x1 | |
| target = x1 - x0 | |
| with torch.autocast("cuda", dtype=torch.bfloat16): | |
| v = model(xt, t, ts, tp) | |
| loss = F.mse_loss(v.float(), target) | |
| opt.zero_grad(set_to_none=True) | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) | |
| opt.step() | |
| d = args.ema if step > args.warmup else 0.0 | |
| with torch.no_grad(): | |
| for pe, pm in zip(ema.parameters(), model.parameters()): | |
| pe.mul_(d).add_(pm.detach(), alpha=1 - d) | |
| for be, bm in zip(ema.buffers(), model.buffers()): | |
| be.copy_(bm) | |
| run_loss += loss.item() | |
| if (step + 1) % args.log_every == 0: | |
| rate = (step + 1 - start) / (time.time() - t0) | |
| print(f"[s{step+1:06d}] loss={run_loss/args.log_every:.4f} lr={lr_at(step):.2e} " | |
| f"{rate:.1f} it/s", flush=True) | |
| run_loss = 0.0 | |
| if (step + 1) % args.ckpt_every == 0: | |
| save(step + 1, "latest") | |
| save(step + 1, f"step{step + 1}") | |
| save(args.steps, "final") | |
| print(f"[train] done in {(time.time()-t0)/60:.1f} min", flush=True) | |
| if __name__ == "__main__": | |
| main() | |