PixelModel-v5 / train_v5.py
TobiasLogic's picture
PixelModel v5: same architecture, 36x the data
a012ee3 verified
Raw
History Blame Contribute Delete
7.9 kB
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 transformers import CLIPTextModel
from dit import DiT
def build_cache(data, cache):
shards = sorted(glob.glob(f"{data}/shard_*.npz"))
if not shards:
raise SystemExit(f"no shards in {data}")
lat, tok = [], []
for i, s in enumerate(shards):
z = np.load(s)
lat.append(z["latents"]); tok.append(z["tokens"])
if (i + 1) % 50 == 0:
print(f" loaded {i+1}/{len(shards)} shards", flush=True)
lat = np.concatenate(lat); tok = np.concatenate(tok)
np.save(f"{cache}_lat.npy", lat); np.save(f"{cache}_tok.npy", tok)
return lat, tok
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--data", default="/root/v5data")
ap.add_argument("--cache", default="/root/v5cache")
ap.add_argument("--out", default="/root/runs/pm5")
ap.add_argument("--steps", type=int, default=120000)
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("--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("--clip", default="openai/clip-vit-base-patch32")
ap.add_argument("--resume", default="")
ap.add_argument("--seed", type=int, default=0)
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")
tok = np.load(f"{args.cache}_tok.npy")
else:
lat, tok = build_cache(args.data, 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)
txt = CLIPTextModel.from_pretrained(args.clip).to(dev).eval()
for p in txt.parameters():
p.requires_grad_(False)
@torch.no_grad()
def encode(ids):
o = txt(input_ids=ids)
return o.last_hidden_state.float(), o.pooler_output.float()
null_ids = torch.full((1, tok.shape[1]), 0, dtype=torch.long, device=dev)
null_ids[0, 0] = 49406; null_ids[0, 1:] = 49407
null_seq, null_pool = encode(null_ids)
tok_t = torch.from_numpy(tok.astype(np.int64))
vlat = torch.from_numpy(np.asarray(lat[val_i])).float()
vseq, vpool = [], []
with torch.no_grad():
for i in range(0, len(val_i), 512):
s, p = encode(tok_t[val_i[i:i+512]].to(dev))
vseq.append(s); vpool.append(p)
vseq = torch.cat(vseq); vpool = torch.cat(vpool)
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))
model = DiT(dim=args.dim, depth=args.depth, heads=args.heads).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] {sum(p.numel() for p in model.parameters())/1e6:.2f}M trainable", 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)
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], 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, t0 = 0.0, time.time()
for step in range(start, args.steps):
i = tr_i[np.random.randint(0, len(tr_i), args.batch)]
x1 = torch.from_numpy(np.asarray(lat[np.sort(i)])).to(dev).float()
seq, pool = encode(tok_t[np.sort(i)].to(dev))
drop = torch.rand(x1.shape[0], device=dev, generator=gen) < args.cfg_dropout
seq = torch.where(drop[:, None, None], null_seq, seq)
pool = torch.where(drop[:, None], null_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
for g in opt.param_groups:
g["lr"] = lr_at(step)
with torch.autocast("cuda", dtype=torch.bfloat16):
v = model(xt, t, seq, pool)
loss = F.mse_loss(v.float(), target)
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()
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} lr={lr_at(step):.2e} "
f"gnorm={gn:.2f} {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,
"steps_per_s": sps}) + "\n"); logf.flush()
run, t0 = 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}, 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}, f"{args.out}/latest.pt")
print("TRAINDONE best_val", best, flush=True)
if __name__ == "__main__":
main()