"""Train the diffusion denoiser or the AR-FIM baseline. Same data, size, schedule. Usage: python -m ml.train --mode diffusion --train data/train.jsonl --out runs/diff python -m ml.train --mode ar --train data/train.jsonl --out runs/ar """ from __future__ import annotations import argparse import json import math import os import time import torch from torch.utils.data import DataLoader from . import ar, diffusion from .config import ModelConfig, TaskConfig, TrainConfig from .data import InfillDataset, load_records from .model import Transformer, amp_ctx from .tokenizer import Tokenizer def pick_device() -> str: if torch.backends.mps.is_available(): return "mps" if torch.cuda.is_available(): return "cuda" return "cpu" def lr_at(step, tc: TrainConfig): """Warmup-stable-decay (TRAINING.md).""" if step < tc.warmup: return tc.lr * step / max(1, tc.warmup) decay_start = int(tc.steps * 0.8) if step < decay_start: return tc.lr frac = (step - decay_start) / max(1, tc.steps - decay_start) return tc.lr * (1.0 - 0.9 * frac) # decay to 0.1*lr def main(): ap = argparse.ArgumentParser() ap.add_argument("--mode", choices=["diffusion", "ar"], required=True) ap.add_argument("--train", default="data/train.jsonl") ap.add_argument("--out", required=True) ap.add_argument("--steps", type=int, default=TrainConfig.steps) ap.add_argument("--batch", type=int, default=TrainConfig.batch_size) ap.add_argument("--d_model", type=int, default=ModelConfig.d_model) ap.add_argument("--layers", type=int, default=ModelConfig.n_layers) ap.add_argument("--heads", type=int, default=ModelConfig.n_heads) ap.add_argument("--ff", type=int, default=ModelConfig.d_ff) ap.add_argument("--tok", choices=["char", "lua"], default="char", help="tokenizer level") ap.add_argument("--seq_len", type=int, default=TaskConfig.seq_len) ap.add_argument("--block_len", type=int, default=TaskConfig.block_len) ap.add_argument("--seed", type=int, default=0) args = ap.parse_args() torch.manual_seed(args.seed) device = pick_device() os.makedirs(args.out, exist_ok=True) records = load_records(args.train) tok = Tokenizer.build([r["source"] for r in records], mode=args.tok) tok.save(os.path.join(args.out, "tokenizer.json")) taskcfg = TaskConfig(seq_len=args.seq_len, block_len=args.block_len) mcfg = ModelConfig( vocab_size=tok.vocab_size, d_model=args.d_model, n_layers=args.layers, n_heads=args.heads, d_ff=args.ff, max_len=taskcfg.seq_len, ) tc = TrainConfig(batch_size=args.batch, steps=args.steps, seed=args.seed) ds = InfillDataset(records, tok, taskcfg, mode=args.mode) print(f"[{args.mode}] device={device} vocab={tok.vocab_size} " f"examples={len(ds)} skipped={ds.skipped} (too long)") dl = DataLoader(ds, batch_size=tc.batch_size, shuffle=True, drop_last=True) model = Transformer(mcfg, causal=(args.mode == "ar")).to(device) print(f"[{args.mode}] params={model.num_params()/1e6:.2f}M") opt = torch.optim.AdamW(model.parameters(), lr=tc.lr, weight_decay=tc.weight_decay) model.train() step = 0 t0 = time.time() running = 0.0 nloss = 0 while step < tc.steps: for batch in dl: if step >= tc.steps: break for g in opt.param_groups: g["lr"] = lr_at(step, tc) with amp_ctx(device): if args.mode == "diffusion": ids, region, block_id, attn_mask = (b.to(device) for b in batch) l = diffusion.loss(model, ids, region, block_id, attn_mask, tok) else: ids, loss_mask = (b.to(device) for b in batch) attn_mask = ids != tok.pad_id l = ar.loss(model, ids, loss_mask, attn_mask, tok) opt.zero_grad() l.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), tc.grad_clip) opt.step() running += l.item() nloss += 1 step += 1 if step % tc.log_every == 0: dt = time.time() - t0 print(f"[{args.mode}] step {step}/{tc.steps} " f"loss {running/nloss:.4f} lr {lr_at(step,tc):.2e} " f"{step/dt:.1f} it/s") running = 0.0 nloss = 0 ckpt = { "model": model.state_dict(), "model_cfg": vars(mcfg), "task_cfg": vars(taskcfg), "mode": args.mode, } torch.save(ckpt, os.path.join(args.out, "model.pt")) with open(os.path.join(args.out, "meta.json"), "w") as f: json.dump({"mode": args.mode, "steps": tc.steps, "params_M": model.num_params() / 1e6}, f, indent=2) print(f"[{args.mode}] saved to {args.out}") if __name__ == "__main__": main()