spec100m / code /train_seq.py
Akahsizrr's picture
squash to reclaim LFS quota
a8f07a3
Raw History Blame Contribute Delete
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
@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()