Spaces:
Sleeping
Sleeping
| """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() | |