| 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 audio_dit import AudioDiT |
|
|
| 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, split, dev): |
| meta = json.load(open(f"{data}/{split}_meta.json")) |
| mel = np.load(f"{data}/{split}_mel.npy") |
| seq = np.load(f"{data}/{split}_text_seq.npy") |
| pool = np.load(f"{data}/{split}_text_pool.npy") |
| ci = np.array(meta["pair_clip_idx"], dtype=np.int64) |
| return (torch.from_numpy(mel).to(dev), torch.from_numpy(seq).to(dev), |
| torch.from_numpy(pool).to(dev), torch.from_numpy(ci).to(dev), meta) |
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--data", default="/root/data") |
| ap.add_argument("--out", default="/root/runs/audio_v1") |
| ap.add_argument("--steps", type=int, default=90000) |
| ap.add_argument("--batch-size", type=int, default=128) |
| 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("--val-every", type=int, default=2000) |
| ap.add_argument("--val-batch", type=int, default=256) |
| ap.add_argument("--log-every", type=int, default=200) |
| ap.add_argument("--save-every", type=int, default=5000) |
| 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 |
|
|
| mel_t, seq_t, pool_t, ci_t, meta = load_split(args.data, "development", dev) |
| vmel_t, vseq_t, vpool_t, vci_t, vmeta = load_split(args.data, "validation", 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() |
|
|
| N = ci_t.shape[0] |
| x_res, y_res = meta["x_res"], meta["y_res"] |
|
|
| def mel_batch(mel_tensor, ci_tensor, idx): |
| return (mel_tensor[ci_tensor[idx]].float() / 127.5 - 1.0).unsqueeze(1) |
|
|
| vg = torch.Generator(device=dev).manual_seed(1234) |
| n_val = min(args.val_batch, vci_t.shape[0]) |
| vidx = torch.arange(n_val, device=dev) |
| vx1 = mel_batch(vmel_t, vci_t, vidx) |
| vs = vseq_t[vidx].float() |
| vp = vpool_t[vidx].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() |
| tb = vt.view(-1, 1, 1, 1) |
| xt = (1 - tb) * vx0 + tb * vx1 |
| with torch.autocast("cuda", dtype=torch.bfloat16): |
| v = ema(xt, vt, vs, vp) |
| return F.mse_loss(v.float(), vx1 - vx0).item() |
|
|
| model = AudioDiT(x_res=x_res, y_res=y_res, |
| text_seq_dim=meta["seq_dim"], text_pool_dim=meta["pool_dim"]).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 |
| best_val = 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_val = ck.get("best_val", float("inf")) |
|
|
| B = args.batch_size |
| gen = torch.Generator(device=dev).manual_seed(args.seed) |
| epochs = args.steps * B / N |
| print(f"[train] {N} train pairs ({meta['n_clips']} clips) + {n_val} val, " |
| f"{model.num_params()/1e6:.2f}M params, grid {model.grid_h}x{model.grid_w} " |
| f"({model.grid_h*model.grid_w} tokens)", flush=True) |
| print(f"[train] {args.steps} steps x batch {B} = {epochs:.0f} epochs, 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): |
| idx = torch.randint(0, N, (B,), device=dev, generator=gen) |
| x1 = mel_batch(mel_t, ci_t, idx) |
| s = seq_t[idx].float() |
| p = pool_t[idx].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) |
| 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() |
| tag = "" |
| if vl < best_val: |
| best_val = vl |
| save_file({k: v.contiguous() for k, v in ema.state_dict().items()}, |
| f"{args.out}/model_best.safetensors") |
| json.dump({"step": step + 1, "val_loss": vl}, open(f"{args.out}/best.json", "w")) |
| tag = " (new best, saved model_best.safetensors)" |
| print(f"[s{step+1:06d}] val_loss={vl:.5f} (ema, {n_val} held out){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_val": best_val}, |
| 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": {"x_res": x_res, "y_res": y_res, "patch": 16, "dim": 384, "depth": 12, |
| "heads": 6, "text_seq_dim": meta["seq_dim"], "text_pool_dim": meta["pool_dim"]}, |
| "mel": {"sample_rate": meta["sample_rate"], "n_fft": meta["n_fft"], |
| "hop_length": meta["hop_length"], "top_db": meta["top_db"]}, |
| "steps": args.steps, "batch_size": args.batch_size, |
| "train_pairs": N, "train_clips": meta["n_clips"], "best_val": best_val}, |
| open(f"{args.out}/config.json", "w"), indent=2) |
| print("TRAINDONE", flush=True) |
|
|
| if __name__ == "__main__": |
| main() |
|
|