File size: 4,991 Bytes
6c9c825
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
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()