| from __future__ import annotations |
|
|
| import argparse |
| import copy |
| import glob |
| import json |
| import math |
| import os |
| import time |
|
|
| import numpy as np |
| import torch |
| import torch.nn.functional as F |
| from diffusers import AutoencoderKL |
| from transformers import AutoModel, CLIPTextModel, CLIPTokenizer, T5EncoderModel, T5TokenizerFast |
|
|
| from dit_v6 import MMDiT |
|
|
| IMAGENET_MEAN = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1) |
| IMAGENET_STD = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1) |
|
|
| def build_cache(shards_dir, cache): |
| shards = sorted(glob.glob(f"{shards_dir}/shard_*.npz")) |
| if not shards: |
| raise SystemExit(f"no shards in {shards_dir}") |
| lat, t5i, t5m, ci = [], [], [], [] |
| for i, s in enumerate(shards): |
| z = np.load(s) |
| lat.append(z["latents"]); t5i.append(z["t5_ids"]); t5m.append(z["t5_mask"]); ci.append(z["clip_ids"]) |
| if (i + 1) % 50 == 0: |
| print(f" loaded {i+1}/{len(shards)} shards", flush=True) |
| lat = np.concatenate(lat); t5i = np.concatenate(t5i); t5m = np.concatenate(t5m); ci = np.concatenate(ci) |
| np.save(f"{cache}_lat.npy", lat) |
| np.save(f"{cache}_t5ids.npy", t5i) |
| np.save(f"{cache}_t5mask.npy", t5m) |
| np.save(f"{cache}_clipids.npy", ci) |
| return lat, t5i, t5m, ci |
|
|
| def repa_weight_at(step, total, warm_frac=0.40, zero_frac=0.70, peak=0.5): |
| p = step / total |
| if p < warm_frac: |
| return peak |
| if p < zero_frac: |
| return peak * (1 - (p - warm_frac) / (zero_frac - warm_frac)) |
| return 0.0 |
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--shards", default="/root/v6cache/shards") |
| ap.add_argument("--cache", default="/root/v6cache/cache") |
| ap.add_argument("--out", default="/root/runs/pm6") |
| ap.add_argument("--steps", type=int, default=150000) |
| ap.add_argument("--batch", type=int, default=192) |
| ap.add_argument("--lr", type=float, default=2e-4) |
| ap.add_argument("--warmup", type=int, default=1500) |
| ap.add_argument("--dim", type=int, default=512) |
| ap.add_argument("--depth", type=int, default=16) |
| ap.add_argument("--heads", type=int, default=8) |
| ap.add_argument("--mlp-hidden", type=int, default=1408) |
| ap.add_argument("--t5-len", type=int, default=32) |
| ap.add_argument("--repa-layer", type=int, default=8) |
| ap.add_argument("--repa-peak", type=float, default=0.5) |
| ap.add_argument("--repa-batch", type=int, default=48) |
| ap.add_argument("--cfg-dropout", type=float, default=0.1) |
| ap.add_argument("--ema", type=float, default=0.9999) |
| ap.add_argument("--val-size", type=int, default=4096) |
| ap.add_argument("--val-every", type=int, default=2000) |
| ap.add_argument("--log-every", type=int, default=200) |
| ap.add_argument("--save-every", type=int, default=5000) |
| ap.add_argument("--vae", default="madebyollin/sdxl-vae-fp16-fix") |
| ap.add_argument("--clip", default="openai/clip-vit-base-patch32") |
| ap.add_argument("--t5", default="google/flan-t5-base") |
| ap.add_argument("--dinov2", default="facebook/dinov2-small") |
| ap.add_argument("--resume", default="") |
| ap.add_argument("--seed", type=int, default=0) |
| ap.add_argument("--grad-ckpt", action="store_true") |
| args = ap.parse_args() |
|
|
| dev = "cuda" |
| os.makedirs(args.out, exist_ok=True) |
| torch.manual_seed(args.seed) |
| torch.backends.cuda.matmul.allow_tf32 = True |
| torch.backends.cudnn.allow_tf32 = True |
|
|
| if os.path.exists(f"{args.cache}_lat.npy"): |
| lat = np.load(f"{args.cache}_lat.npy", mmap_mode="r") |
| t5i = np.load(f"{args.cache}_t5ids.npy") |
| t5m = np.load(f"{args.cache}_t5mask.npy") |
| ci = np.load(f"{args.cache}_clipids.npy") |
| else: |
| lat, t5i, t5m, ci = build_cache(args.shards, args.cache) |
|
|
| N = len(lat) |
| perm = np.random.RandomState(args.seed).permutation(N) |
| val_i = np.sort(perm[:args.val_size]) |
| tr_i = perm[args.val_size:] |
| print(f"[data] {N} pairs, {len(tr_i)} train, {len(val_i)} val", flush=True) |
|
|
| vae = AutoencoderKL.from_pretrained(args.vae).to(dev).half().eval() |
| vae_scale = vae.config.scaling_factor |
| for p in vae.parameters(): |
| p.requires_grad_(False) |
|
|
| clip_txt = CLIPTextModel.from_pretrained(args.clip).to(dev).eval() |
| for p in clip_txt.parameters(): |
| p.requires_grad_(False) |
|
|
| t5 = T5EncoderModel.from_pretrained(args.t5).to(dev).eval() |
| for p in t5.parameters(): |
| p.requires_grad_(False) |
|
|
| dinov2 = AutoModel.from_pretrained(args.dinov2).to(dev).eval() |
| for p in dinov2.parameters(): |
| p.requires_grad_(False) |
| imagenet_mean = IMAGENET_MEAN.to(dev) |
| imagenet_std = IMAGENET_STD.to(dev) |
|
|
| @torch.no_grad() |
| def encode_text(t5_ids, t5_mask, clip_ids): |
| t5_out = t5(input_ids=t5_ids, attention_mask=t5_mask).last_hidden_state.float() |
| clip_pool = clip_txt(input_ids=clip_ids).pooler_output.float() |
| return t5_out, clip_pool |
|
|
| @torch.no_grad() |
| def dino_features(x1_latent): |
| px = vae.decode((x1_latent / vae_scale).to(vae.dtype)).sample.float() |
| px = (px.clamp(-1, 1) + 1) / 2 |
| px = F.interpolate(px, size=(224, 224), mode="bilinear", align_corners=False) |
| px = (px - imagenet_mean) / imagenet_std |
| out = dinov2(pixel_values=px.to(dinov2.dtype)).last_hidden_state |
| return out[:, 1:, :].float() |
|
|
| t5_tok = T5TokenizerFast.from_pretrained(args.t5) |
| clip_tok = CLIPTokenizer.from_pretrained(args.clip) |
| null_t5_enc = t5_tok([""], padding="max_length", max_length=args.t5_len, truncation=True, return_tensors="pt") |
| null_t5_ids = null_t5_enc["input_ids"].to(dev) |
| null_t5_mask = null_t5_enc["attention_mask"].to(dev) |
| null_clip_ids = clip_tok([""], padding="max_length", max_length=ci.shape[1], truncation=True, |
| return_tensors="pt")["input_ids"].to(dev) |
| null_t5_seq, null_clip_pool = encode_text(null_t5_ids, null_t5_mask, null_clip_ids) |
|
|
| t5i_t = torch.from_numpy(t5i.astype(np.int64)) |
| t5m_t = torch.from_numpy(t5m.astype(np.int64)) |
| ci_t = torch.from_numpy(ci.astype(np.int64)) |
|
|
| vlat = torch.from_numpy(np.asarray(lat[val_i])).float() |
| vg = torch.Generator(device=dev).manual_seed(1234) |
| vx1 = vlat.to(dev) |
| vx0 = torch.randn(vx1.shape, device=dev, generator=vg) |
| vt = torch.sigmoid(torch.randn(vx1.shape[0], device=dev, generator=vg)) |
| with torch.no_grad(): |
| vseq, vmask, vpool = [], [], [] |
| for i in range(0, len(val_i), 256): |
| s, p = encode_text(t5i_t[val_i[i:i+256]].to(dev), t5m_t[val_i[i:i+256]].to(dev), ci_t[val_i[i:i+256]].to(dev)) |
| vseq.append(s); vmask.append(t5m_t[val_i[i:i+256]].to(dev)); vpool.append(p) |
| vseq = torch.cat(vseq); vmask = torch.cat(vmask); vpool = torch.cat(vpool) |
|
|
| model = MMDiT(dim=args.dim, depth=args.depth, heads=args.heads, mlp_hidden=args.mlp_hidden, |
| t5_len=args.t5_len).to(dev) |
| 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) |
| print(f"[model] {model.num_params()/1e6:.2f}M trainable, {model.num_backbone_params()/1e6:.2f}M backbone", flush=True) |
|
|
| start, best = 0, float("inf") |
| 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"] + 1; best = ck.get("best", best) |
| print(f"[resume] from step {start}", flush=True) |
|
|
| def lr_at(s): |
| if s < args.warmup: |
| return args.lr * (s + 1) / args.warmup |
| p = (s - 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)))) |
|
|
| @torch.no_grad() |
| def val_loss(): |
| tot, n = 0.0, 0 |
| for j in range(0, vx1.shape[0], args.batch): |
| sl = slice(j, j + args.batch) |
| m = vx1[sl].shape[0] |
| tb = vt[sl].view(-1, 1, 1, 1) |
| xt = (1 - tb) * vx0[sl] + tb * vx1[sl] |
| with torch.autocast("cuda", dtype=torch.bfloat16): |
| v = ema(xt, vt[sl], vseq[sl], vmask[sl], vpool[sl]) |
| tot += F.mse_loss(v.float(), vx1[sl] - vx0[sl]).item() * m |
| n += m |
| return tot / n |
|
|
| logf = open(f"{args.out}/log.jsonl", "a") |
| gen = torch.Generator(device=dev).manual_seed(args.seed) |
| run, run_diff, run_repa = 0.0, 0.0, 0.0 |
| t0 = time.time() |
| for step in range(start, args.steps): |
| i = tr_i[np.random.randint(0, len(tr_i), args.batch)] |
| i_sorted = np.sort(i) |
| x1 = torch.from_numpy(np.asarray(lat[i_sorted])).to(dev).float() |
| t5_ids_b = t5i_t[i_sorted].to(dev) |
| t5_mask_b = t5m_t[i_sorted].to(dev) |
| clip_ids_b = ci_t[i_sorted].to(dev) |
| seq, pool = encode_text(t5_ids_b, t5_mask_b, clip_ids_b) |
| mask = t5_mask_b |
|
|
| drop = torch.rand(x1.shape[0], device=dev, generator=gen) < args.cfg_dropout |
| seq = torch.where(drop[:, None, None], null_t5_seq, seq) |
| mask = torch.where(drop[:, None], null_t5_mask, mask) |
| pool = torch.where(drop[:, None], null_clip_pool, pool) |
|
|
| x0 = torch.randn(x1.shape, device=dev, generator=gen) |
| t = torch.sigmoid(torch.randn(x1.shape[0], device=dev, generator=gen)) |
| tb = t.view(-1, 1, 1, 1) |
| xt = (1 - tb) * x0 + tb * x1 |
| target = x1 - x0 |
|
|
| rw = repa_weight_at(step, args.steps, peak=args.repa_peak) |
|
|
| for g in opt.param_groups: |
| g["lr"] = lr_at(step) |
|
|
| with torch.autocast("cuda", dtype=torch.bfloat16): |
| if rw > 0: |
| v, repa_pred = model(xt, t, seq, mask, pool, return_repa=True, use_checkpoint=args.grad_ckpt) |
| else: |
| v = model(xt, t, seq, mask, pool, use_checkpoint=args.grad_ckpt) |
| loss_diff = F.mse_loss(v.float(), target) |
| if rw > 0: |
| rb = min(args.repa_batch, x1.shape[0]) |
| with torch.no_grad(): |
| dino_tgt = dino_features(x1[:rb]) |
| loss_repa = 1.0 - F.cosine_similarity(repa_pred[:rb].float(), dino_tgt, dim=-1).mean() |
| loss = loss_diff + rw * loss_repa |
| else: |
| loss_repa = torch.zeros((), device=dev) |
| loss = loss_diff |
|
|
| opt.zero_grad(set_to_none=True) |
| loss.backward() |
| gn = 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.item(); run_diff += loss_diff.item(); run_repa += loss_repa.item() |
| if (step + 1) % args.log_every == 0: |
| el = time.time() - t0 |
| sps = args.log_every / el |
| print(f"[s{step+1:06d}] loss={run/args.log_every:.4f} diff={run_diff/args.log_every:.4f} " |
| f"repa={run_repa/args.log_every:.4f} rw={rw:.3f} lr={lr_at(step):.2e} gnorm={gn:.2f} " |
| f"{sps:.2f} steps/s eta={(args.steps-step-1)/sps/3600:.1f}h", flush=True) |
| logf.write(json.dumps({"step": step + 1, "loss": run/args.log_every, "loss_diff": run_diff/args.log_every, |
| "loss_repa": run_repa/args.log_every, "repa_weight": rw, |
| "steps_per_s": sps}) + "\n"); logf.flush() |
| run, run_diff, run_repa, t0 = 0.0, 0.0, 0.0, time.time() |
|
|
| if (step + 1) % args.val_every == 0 or step + 1 == args.steps: |
| vl = val_loss() |
| tag = "" |
| if vl < best: |
| best = vl |
| torch.save({"ema": ema.state_dict(), "step": step, "val": vl, |
| "cfg": {"dim": args.dim, "depth": args.depth, "heads": args.heads, |
| "mlp_hidden": args.mlp_hidden, "t5_len": args.t5_len}}, |
| f"{args.out}/best.pt") |
| tag = " *best*" |
| print(f"[s{step+1:06d}] val_loss={vl:.5f}{tag}", flush=True) |
| logf.write(json.dumps({"step": step + 1, "val_loss": vl}) + "\n"); logf.flush() |
| t0 = time.time() |
|
|
| if (step + 1) % args.save_every == 0 or step + 1 == args.steps: |
| torch.save({"model": model.state_dict(), "ema": ema.state_dict(), "opt": opt.state_dict(), |
| "step": step, "best": best, |
| "cfg": {"dim": args.dim, "depth": args.depth, "heads": args.heads, |
| "mlp_hidden": args.mlp_hidden, "t5_len": args.t5_len}}, |
| f"{args.out}/latest.pt") |
|
|
| print("TRAINDONE best_val", best, flush=True) |
|
|
| if __name__ == "__main__": |
| main() |
|
|