VoxelModel-v1 / train.py
TobiasLogic's picture
VoxelModel v1: tiny text-to-3D voxel diffusion
49ce4bd verified
Raw
History Blame Contribute Delete
7.7 kB
from __future__ import annotations
import argparse
import copy
import json
import math
import os
import time
import numpy as np
import torch
import torch.nn.functional as F
from safetensors.torch import save_file
from voxel_dit import VoxelDiT
def cosine_lr(step, total, base, floor, warmup):
if step < warmup:
return base * (step + 1) / max(1, warmup)
t = (step - warmup) / max(1, total - warmup)
return floor + 0.5 * (base - floor) * (1 + math.cos(math.pi * t))
def load_split(data, filtered, n_val, seed=0):
meta = json.load(open(f"{data}/meta.json"))
keep = np.array(meta["keep"], bool)
idx = np.nonzero(keep)[0] if filtered else np.arange(len(keep))
perm = np.random.RandomState(seed).permutation(len(idx))
idx = idx[perm]
val, tr = idx[:n_val], idx[n_val:]
packed = np.load(f"{data}/voxels_packed.npy")
seq = np.load(f"{data}/text_seq.npy")
pool = np.load(f"{data}/text_pool.npy")
pack = lambda s: (np.unpackbits(packed[s], axis=1), seq[s], pool[s])
return pack(tr), pack(val), len(keep)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--data", default="/root/data")
ap.add_argument("--out", default="/root/runs/voxel_v1")
ap.add_argument("--steps", type=int, default=90000)
ap.add_argument("--batch-size", type=int, default=256)
ap.add_argument("--lr", type=float, default=2e-4)
ap.add_argument("--min-lr", type=float, default=1e-6)
ap.add_argument("--warmup", type=int, default=500)
ap.add_argument("--cfg-dropout", type=float, default=0.1)
ap.add_argument("--ema", type=float, default=0.9999)
ap.add_argument("--filtered", type=int, default=1)
ap.add_argument("--val-size", type=int, default=1024)
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=10000)
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
(vox, seq, pool), (vvox, vseq, vpool), total_meshes = load_split(
args.data, args.filtered, args.val_size)
M = vox.shape[0]
vox_t = torch.from_numpy(vox).to(dev)
seq_t = torch.from_numpy(seq).to(dev)
pool_t = torch.from_numpy(pool).to(dev)
null_seq = torch.from_numpy(np.load(f"{args.data}/null_seq.npy")).to(dev).float()
null_pool = torch.from_numpy(np.load(f"{args.data}/null_pool.npy")).to(dev).float()
vg = torch.Generator(device=dev).manual_seed(1234)
vx1 = torch.from_numpy(vvox).to(dev).view(-1, 1, 32, 32, 32).float() * 2 - 1
vs = torch.from_numpy(vseq).to(dev).float()
vp = torch.from_numpy(vpool).to(dev).float()
vx0 = torch.randn(vx1.shape, device=dev, generator=vg)
vt = torch.sigmoid(torch.randn(vx1.shape[0], device=dev, generator=vg))
@torch.no_grad()
def val_loss():
ema.eval()
tot, n = 0.0, 0
for j in range(0, vx1.shape[0], args.batch_size):
sl = slice(j, j + args.batch_size)
m = vx1[sl].shape[0]
tb = vt[sl].view(-1, 1, 1, 1, 1)
xt = (1 - tb) * vx0[sl] + tb * vx1[sl]
with torch.autocast("cuda", dtype=torch.bfloat16):
v = ema(xt, vt[sl], vs[sl], vp[sl])
tot += F.mse_loss(v.float(), vx1[sl] - vx0[sl]).item() * m
n += m
return tot / n
model = VoxelDiT().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)
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"] + 1
B = args.batch_size
gen = torch.Generator(device=dev).manual_seed(args.seed)
epochs = args.steps * B / M
print(f"[train] {M} train + {vx1.shape[0]} val of {total_meshes} captioned "
f"({'filtered' if args.filtered else 'unfiltered'}), "
f"{model.num_params()/1e6:.2f}M params", flush=True)
print(f"[train] {args.steps} steps x batch {B} = {epochs:.0f} epochs, "
f"resume from {start}", flush=True)
logf = open(f"{args.out}/log.jsonl", "a")
running, t0 = 0.0, time.time()
for step in range(start, args.steps):
i = torch.randint(0, M, (B,), device=dev, generator=gen)
x1 = vox_t[i].view(B, 1, 32, 32, 32).float() * 2 - 1
s = seq_t[i].float()
p = pool_t[i].float()
drop = torch.rand(B, device=dev, generator=gen) < args.cfg_dropout
s = torch.where(drop[:, None, None], null_seq, s)
p = torch.where(drop[:, None], null_pool, p)
x0 = torch.randn(x1.shape, device=dev, generator=gen)
u = torch.randn(B, device=dev, generator=gen)
t = torch.sigmoid(u)
tb = t.view(-1, 1, 1, 1, 1)
xt = (1 - tb) * x0 + tb * x1
target = x1 - x0
lr = cosine_lr(step, args.steps, args.lr, args.min_lr, args.warmup)
for g in opt.param_groups:
g["lr"] = lr
with torch.autocast("cuda", dtype=torch.bfloat16):
v = model(xt, t, s, p)
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()
with torch.no_grad():
d = 1 - args.ema
for pe, pm in zip(ema.parameters(), model.parameters()):
pe.mul_(args.ema).add_(pm.detach(), alpha=d)
running += loss.item()
if (step + 1) % args.log_every == 0:
avg = running / args.log_every
el = time.time() - t0
sps = args.log_every / el
eta = (args.steps - step - 1) / sps / 3600
print(f"[s{step+1:06d}] loss={avg:.5f} lr={lr:.2e} gnorm={gn:.2f} "
f"{sps:.2f} steps/s eta={eta:.2f}h", flush=True)
logf.write(json.dumps({"step": step + 1, "loss": avg, "lr": lr,
"gnorm": float(gn), "steps_per_s": sps}) + "\n")
logf.flush()
running, t0 = 0.0, time.time()
if (step + 1) % args.val_every == 0 or step + 1 == args.steps:
vl = val_loss()
print(f"[s{step+1:06d}] val_loss={vl:.5f} (ema, {args.val_size} held out)", 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},
f"{args.out}/ckpt.pt")
save_file({k: v.contiguous() for k, v in ema.state_dict().items()},
f"{args.out}/model.safetensors")
print(f"[s{step+1:06d}] saved", flush=True)
json.dump({"dit": {"vox_size": 32, "patch": 4, "dim": 384, "depth": 12,
"heads": 6, "text_dim": 512},
"steps": args.steps, "batch_size": args.batch_size,
"meshes": M, "filtered": bool(args.filtered)},
open(f"{args.out}/config.json", "w"), indent=2)
print("TRAINDONE", flush=True)
if __name__ == "__main__":
main()