Download code/train_seq.py from Akahsizrr/spec100m: direct link, hf CLI and curl.
- Browser
- Download file 11.3 kB
-
https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/train_seq.py
- Command line
-
hf download hf://Akahsizrr/spec100m/code/train_seq.py
-
curl -L -o train_seq.py https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/train_seq.py
11.3 kB
| """Train the EAGLE-lite sequential draft module on the frozen base. | |
| The seq module's job: given h_t and the next committed token, evolve a | |
| hidden state h~ that (a) predicts the following token AND (b) stays close | |
| to the base's own hidden state — the feature regression that lets the | |
| draft chain survive its own mistakes. | |
| Loss: L = CE(h~ @ E^T, next_token_or_base_argmax) + reg_w * SmoothL1(h~, h_next) | |
| One backward trains all K chain positions at once — far more sample- | |
| efficient than the flat heads' per-head sampling. | |
| Usage: | |
| python train_seq.py --resume <ckpt> --steps 1000 --batch 16 --seq 512 \ | |
| --chain 32 --reg-w 0.5 | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import os | |
| import time | |
| import torch | |
| import torch.nn.functional as F | |
| from config import Config | |
| from model import build_model | |
| from data_pipeline import load_tokens | |
| from train import batched, TrainingController | |
| from train_medusa import eval_streak | |
| def eval_seq_streak(model, cfg, data, seq=1024, n_anchors=256, | |
| device="cuda", chain=32): | |
| """Free-running seq-draft streak vs base greedy — same metric as | |
| eval_streak but through the sequential module.""" | |
| E = model.tok_emb.weight | |
| start = torch.randint(0, data.numel() - seq - 1, (1,), device=device) | |
| ids = data[start:start + seq].unsqueeze(0) | |
| h = model(ids) | |
| base_pred = model.lm_head(h).argmax(-1)[0] | |
| lo, hi = cfg.medusa_cond_group - 1, seq - chain - 3 | |
| anchors = torch.randint(lo, hi, (n_anchors,), device=device).sort().values | |
| h_a = h[0, anchors] | |
| # chain input token = the REAL next token (teacher-forced start, then | |
| # self-fed) — mirrors inference where pending = committed token | |
| streaks = [] | |
| acc1 = [] | |
| B = 64 | |
| for s in range(0, n_anchors, B): | |
| aa = anchors[s:s + B] | |
| h_sub = h_a[s:s + B] | |
| # pending = base's own argmax at anchor (the committed token) | |
| pend = (h_sub @ E.T).argmax(-1, keepdim=True) | |
| G = cfg.medusa_cond_group | |
| hist = torch.stack([ids[0, t - G + 1:t + 1] for t in aa]) | |
| draft = model.spec_draft_seq(h_sub, prefix_ids=pend, | |
| hist_ids=hist, n=chain) | |
| # draft[i] is the token for position anchor+2+i; base_pred[anchor+1+i] | |
| # is base's greedy for that position | |
| tgt = torch.stack([base_pred[t + 1:t + 1 + chain] for t in aa]) | |
| match = draft == tgt | |
| st = match.cumprod(1).sum(1).float() | |
| streaks.append(st) | |
| acc1.append(match[:, 0].float()) | |
| return (torch.cat(streaks).mean().item(), | |
| torch.cat(acc1).mean().item()) | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--data", type=str, default="mixture500m") | |
| ap.add_argument("--steps", type=int, default=1000) | |
| ap.add_argument("--seq", type=int, default=512) | |
| ap.add_argument("--batch", type=int, default=16) | |
| ap.add_argument("--chain", type=int, default=32, | |
| help="draft chain length (K positions trained per anchor)") | |
| ap.add_argument("--anchors", type=int, default=8, | |
| help="anchor windows per sequence per step") | |
| ap.add_argument("--lr", type=float, default=5e-4) | |
| ap.add_argument("--warmup", type=int, default=50) | |
| ap.add_argument("--reg-w", type=float, default=0.5, | |
| help="feature-regression weight: ||h~ - h_base_next||") | |
| ap.add_argument("--base-label", type=float, default=1.0, | |
| help="fraction of steps using base argmax as CE target") | |
| ap.add_argument("--ss-prob", type=float, default=0.0, | |
| help="scheduled sampling: fraction of chain-input " | |
| "positions fed the module's OWN argmax instead of " | |
| "the real token (closes the teacher-forcing gap " | |
| "that kills free-running streaks)") | |
| ap.add_argument("--out", type=str, default="checkpoints/spec_seq") | |
| ap.add_argument("--eval_every", type=int, default=100) | |
| ap.add_argument("--checkpoint_every", type=int, default=500) | |
| ap.add_argument("--resume", type=str, required=True) | |
| ap.add_argument("--grad-clip", type=float, default=1.0) | |
| args = ap.parse_args() | |
| device = "cuda" | |
| torch.backends.cuda.matmul.allow_tf32 = True | |
| torch.backends.cudnn.allow_tf32 = True | |
| cfg = Config.v5_500m() | |
| assert cfg.medusa_seq_len > 0, "seq draft requires medusa_seq_len > 0" | |
| model = build_model(cfg, device) | |
| ckpt = torch.load(args.resume, map_location=device) | |
| msd = model.state_dict() | |
| # keep only keys whose shape matches — handles arch changes like the | |
| # seq_fc widening (2048 -> 2304 inputs) across checkpoints | |
| sd = {k: v for k, v in ckpt["model"].items() | |
| if k in msd and msd[k].shape == v.shape} | |
| missing = [k for k in msd if k not in sd] | |
| model.load_state_dict(sd, strict=False) | |
| print(f"resumed: {args.resume} (step {ckpt.get('step')}, " | |
| f"loss {ckpt.get('loss'):.4f})") | |
| print(f" missing: {sorted(missing)}") | |
| for name, p in model.named_parameters(): | |
| p.requires_grad = name.startswith("spec_seq") | |
| trainable = [p for p in model.parameters() if p.requires_grad] | |
| print(f"trainable: {sum(p.numel() for p in trainable)/1e6:.1f}M " | |
| f"seq-draft params") | |
| data = load_tokens(args.data).to(device) | |
| opt = torch.optim.AdamW(trainable, lr=args.lr, betas=(0.9, 0.95), | |
| weight_decay=0.01, fused=True) | |
| ctrl = TrainingController(args.lr, args.warmup, args.steps) | |
| os.makedirs(args.out, exist_ok=True) | |
| B, T, K = args.batch, args.seq, args.chain | |
| print(f"\n=== Seq-draft training ===") | |
| print(f" steps {args.steps} batch {B}x{T} chain {K} " | |
| f"anchors/seq {args.anchors} reg_w {args.reg_w}\n") | |
| model.train() | |
| loader = batched(data, T, B, device) | |
| t_start = time.perf_counter() | |
| best_streak = -1.0 | |
| E = model.tok_emb.weight | |
| for step in range(args.steps): | |
| ids = next(loader) | |
| opt.zero_grad(set_to_none=True) | |
| with torch.no_grad(): | |
| h = model(ids) # [B,T,d] | |
| base_am = model.lm_head(h).argmax(-1) # [B,T] | |
| # pick random anchor positions with room for the chain + targets | |
| G = cfg.medusa_cond_group | |
| lo, hi = G - 1, T - K - 2 | |
| A = args.anchors | |
| pos = torch.randint(lo, hi, (B, A), device=device) # [B,A] | |
| offs = torch.arange(K, device=device) | |
| posK = (pos.unsqueeze(-1) + offs).reshape(B, A * K) # [B,A*K] | |
| d = h.shape[-1] | |
| def g(x, o): | |
| return x.gather(1, (posK + o) | |
| .unsqueeze(-1).expand(-1, -1, d) | |
| if x.dim() == 3 else posK + o) | |
| # chain inputs: x_i = [h[t+i] ; emb(tok_{t+i+1}) ; cond(window)] | |
| h_in = g(h, 0).view(B, A, K, d) | |
| h_next = g(h, 1).view(B, A, K, d) | |
| if torch.rand(()) < args.base_label: | |
| tgt = base_am.gather(1, posK + 1).view(B, A, K) | |
| else: | |
| tgt = ids.gather(1, posK + 2).view(B, A, K) | |
| # token span covering [context G ; chain inputs] per anchor: | |
| # span[j] = token at pos-G+1+j (j=0..K+G-1) | |
| offs_span = torch.arange(K + G, device=device) | |
| span_idx = (pos.unsqueeze(-1) - G + 1 + offs_span).clamp(min=0) | |
| span = ids.gather(1, span_idx.reshape(B, -1)).view(B, A, K + G) | |
| if args.ss_prob > 0: | |
| # scheduled sampling: no-grad pass for the chain's own argmax, | |
| # then mix into the input span (pred for input pos i comes | |
| # from chain step i-1) | |
| with torch.no_grad(): | |
| cond0 = (model.tok_emb( | |
| span.unfold(2, G, 1)[:, :, 1:K + 1]) | |
| x0 = model.spec_seq_fc(torch.cat( | |
| [h_in, model.tok_emb(span[:, :, G:G + K]), cond0], -1) | |
| ).reshape(B * A, K, -1) | |
| preds0 = (model.spec_seq_blk(x0) @ E.T).argmax(-1) | |
| preds0 = preds0.view(B, A, K) | |
| mix = span.clone() | |
| ss_mask = torch.rand(B, A, K, device=device) < args.ss_prob | |
| ss_mask[:, :, 0] = False # first input is the real pending tok | |
| mix[:, :, G + 1:G + K] = torch.where( | |
| ss_mask[:, :, 1:], preds0[:, :, :-1], | |
| span[:, :, G + 1:G + K]) | |
| span = mix | |
| tok_in = span[:, :, G:G + K] # [B,A,K] | |
| win = span.unfold(2, G, 1)[:, :, 1:K + 1] # [B,A,K,G] | |
| cond = (model.tok_emb(win) @ model.spec_tok_proj | |
| ).reshape(B, A, K, G * cfg.medusa_emb_rank) | |
| x_in = torch.cat([h_in, model.tok_emb(tok_in), cond], dim=-1) | |
| x = model.spec_seq_fc(x_in).reshape(B * A, K, -1) # [B*A,K,d] | |
| h_out = model.spec_seq_blk(x) # [B*A,K,d] | |
| lg = h_out @ E.T # [B*A,K,V] | |
| l_cls = F.cross_entropy(lg.reshape(-1, lg.shape[-1]).float(), | |
| tgt.reshape(-1)) | |
| l_reg = F.smooth_l1_loss(h_out, h_next.reshape(B * A, K, -1)) | |
| l = l_cls + args.reg_w * l_reg | |
| l.backward() | |
| total = l_cls.item() | |
| grad_norm = torch.nn.utils.clip_grad_norm_(trainable, | |
| args.grad_clip) | |
| opt.step() | |
| if step % 20 == 0 or step == args.steps - 1: | |
| lr, should_ckpt, msg = ctrl.update(total, grad_norm.item(), step) | |
| for pg in opt.param_groups: | |
| pg["lr"] = lr | |
| print(f"step {step:5d} lr {lr:.2e} loss {total:.4f} " | |
| f"grad {grad_norm.item():.2f} " | |
| f"{time.perf_counter()-t_start:.0f}s{msg}", flush=True) | |
| else: | |
| lr = ctrl.lr_at(step + 1) | |
| for pg in opt.param_groups: | |
| pg["lr"] = lr | |
| if (step + 1) % args.eval_every == 0: | |
| model.eval() | |
| streak, acc1 = eval_seq_streak( | |
| model, cfg, data, seq=T, device=device, chain=K) | |
| model.train() | |
| tag = "" | |
| if streak > best_streak: | |
| best_streak = streak | |
| torch.save({"model": model.state_dict(), | |
| "cfg": cfg.__dict__, "step": step + 1, | |
| "loss": total, "streak": streak}, | |
| os.path.join(args.out, "best.pt")) | |
| tag = " (new best streak)" | |
| print(f" [eval] seq streak ~{streak:.1f} " | |
| f"seq-0 acc {acc1:.2f}{tag}", flush=True) | |
| if (step + 1) % args.checkpoint_every == 0: | |
| torch.save({"model": model.state_dict(), "cfg": cfg.__dict__, | |
| "step": step + 1, "loss": total}, | |
| os.path.join(args.out, f"step_{step+1}.pt")) | |
| torch.save({"model": model.state_dict(), "cfg": cfg.__dict__, | |
| "step": args.steps, "loss": total}, | |
| os.path.join(args.out, "final.pt")) | |
| streak, acc1 = eval_seq_streak(model, cfg, data, seq=T, device=device, | |
| chain=K) | |
| print(f"\n=== Done === final seq streak ~{streak:.1f} " | |
| f"seq-0 acc {acc1:.2f}") | |
| if __name__ == "__main__": | |
| main() | |