"""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 --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 @torch.no_grad() 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]) @ model.spec_tok_proj).reshape(B, A, 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()