spec100m / code /inference.py
Akahsizrr's picture
squash to reclaim LFS quota
a8f07a3
Raw History Blame Contribute Delete
42 kB
"""Fast inference engine: KV cache + accept-all speculative decoding.
Two generation modes:
- autoregressive(): 1 token per forward pass (baseline).
- speculative(): K+1 tokens per forward pass (Medusa heads, accept-all).
"Accept everything" = no rejection sampling. We take every draft token the
Medusa heads propose. Steady state is 1 forward pass -> K+1 output tokens,
so the speedup over autoregressive approaches K+1x (minus overhead).
torch.compile is applied to the forward pass if enabled; we gracefully fall
back to eager if compilation fails (dynamic cache length can trip guards).
"""
from __future__ import annotations
import time
import torch
from config import Config
from model import SpecModel
class InferenceEngine:
def __init__(self, model: SpecModel, device="cuda"):
self.model = model
self.cfg = model.cfg
self.device = device
self.dtype = model.tok_emb.weight.dtype
self._cache_k = None
self._cache_v = None
self._compiled = False
# ------------------------------------------------------------------ cache
def alloc_cache(self, batch=1):
cfg = self.cfg
dt = self.dtype
self._cache_k = [
torch.zeros(batch, cfg.max_seq_len, cfg.n_kv_heads, cfg.head_dim,
device=self.device, dtype=dt)
for _ in range(cfg.n_layers)
]
self._cache_v = [t.clone() for t in self._cache_k]
def reset_cache(self):
if self._cache_k is not None:
for t in self._cache_k:
t.zero_()
for t in self._cache_v:
t.zero_()
# ------------------------------------------------------------------ compile
def try_compile(self):
"""Best-effort torch.compile of the forward.
We compile a CLOSURE that captures the KV cache as free variables (not
function inputs). torch.compile forbids in-place mutation of *inputs*,
but in-place writes to captured state (buffers) are allowed. This is the
trick that lets the per-step cache update `cache[:, pos:pos+T] = k` survive
compilation.
"""
if self._compiled:
return True
if self._cache_k is None:
self.alloc_cache()
cache_k, cache_v = self._cache_k, self._cache_v
model = self.model
mode = self.cfg.compile_mode
@torch.compile(mode=mode, dynamic=True)
def _compiled_step(input_ids, start_pos):
return model(input_ids, cache_k, cache_v, start_pos)
self._compiled_step = _compiled_step
self._compiled = True
return True
# ------------------------------------------------------------------ core step
@torch.no_grad()
def step(self, input_ids: torch.Tensor, start_pos: int) -> torch.Tensor:
"""One forward pass over input_ids [B, T] at start_pos. Returns hidden [B, T, d]."""
if self._compiled:
return self._compiled_step(input_ids, start_pos)
return self.model(input_ids, self._cache_k, self._cache_v, start_pos)
@torch.no_grad()
def next_tokens_from_hidden(self, h_last: torch.Tensor):
"""Given hidden state at the last position, return (main_token, medusa_tokens).
main_token: [B] long. medusa_tokens: [B, K] long. All greedy argmax."""
main = self.model.lm_head(h_last).argmax(-1) # [B]
medusa = self.model.medusa_argmax(h_last) # [B, K]
return main, medusa
# ------------------------------------------------------------------ prefill
@torch.no_grad()
def prefill(self, prompt_ids: list[int]) -> torch.Tensor:
"""Process the prompt, fill KV cache, return hidden state at last prompt position."""
ids = torch.tensor([prompt_ids], device=self.device, dtype=torch.long)
h = self.step(ids, start_pos=0) # [1, P, d]
return h[:, -1] # [1, d]
# ------------------------------------------------------------------ AR baseline
@torch.no_grad()
def autoregressive(self, prompt_ids: list[int], new_tokens: int) -> list[int]:
"""1 token per forward pass. Returns generated token ids (excludes prompt)."""
self.reset_cache()
h_last = self.prefill(prompt_ids) # hidden at last prompt pos
pos = len(prompt_ids)
out = []
while len(out) < new_tokens:
tok = self.model.lm_head(h_last).argmax(-1) # [1]
out.append(int(tok.item()))
h = self.step(tok.unsqueeze(0), start_pos=pos) # [1,1,d]
h_last = h[:, -1]
pos += 1
return out[:new_tokens]
# ------------------------------------------------------------------ speculative (accept-all)
@torch.no_grad()
def speculative(self, prompt_ids: list[int], new_tokens: int) -> list[int]:
"""K+1 tokens per forward pass, all accepted (no rejection sampling).
Returns generated token ids (excludes prompt)."""
self.reset_cache()
cfg = self.cfg
K = cfg.medusa_heads
G = cfg.medusa_cond_group if cfg.medusa_rank > 1 else 0
h_last = self.prefill(prompt_ids) # hidden at last prompt pos
pos = len(prompt_ids)
committed = list(prompt_ids)
out = []
if cfg.medusa_rank == 1:
main, medusa = self.next_tokens_from_hidden(h_last)
accepted = torch.cat([main, medusa[0]]) # [K+1]
else:
hist = committed[-G:]
hist = [committed[0]] * (G - len(hist)) + hist
hist_t = torch.tensor([hist], device=self.device)
main = self.model.lm_head(h_last).argmax(-1) # [1]
draft = self.model.spec_draft(h_last, hist_t, first_head=0)
accepted = torch.cat([main, draft[0]]) # [K+1]
out.extend(accepted.tolist())
committed.extend(accepted.tolist())
while len(out) < new_tokens:
h = self.step(accepted.unsqueeze(0), start_pos=pos) # [1, K+1, d]
pos += K + 1
h_last = h[:, -1]
if cfg.medusa_rank == 1:
main, medusa = self.next_tokens_from_hidden(h_last)
accepted = torch.cat([main, medusa[0]])
else:
hist = committed[-G:]
hist_t = torch.tensor([hist], device=self.device)
main = self.model.lm_head(h_last).argmax(-1)
draft = self.model.spec_draft(h_last, hist_t, first_head=0)
accepted = torch.cat([main, draft[0]])
out.extend(accepted.tolist())
committed.extend(accepted.tolist())
return out[:new_tokens]
# ------------------------------------------------------------------ speculative (verified)
@torch.no_grad()
def speculative_verified(self, prompt_ids: list[int], new_tokens: int,
verify_k: int = None):
"""Greedy-verified speculative decoding.
Each round: heads draft K tokens; one causal forward verifies them all;
commit longest prefix matching base argmax + 1 free correction token.
Output is identical to base-model greedy decoding. Quality loss: zero.
verify_k: cap on draft length per round (default = all K heads).
Smaller verify_k = cheaper rounds; optimal when accept streak is short.
Returns (token_ids, stats dict)."""
cfg = self.cfg
assert cfg.medusa_rank > 1, "verified mode needs v6 heads"
G = cfg.medusa_cond_group
Kc, Kp = cfg.medusa_cond_heads, cfg.medusa_par_heads
if verify_k is None or verify_k >= Kc + Kp:
n_cond, n_par = Kc, Kp
elif verify_k >= Kc:
n_cond, n_par = Kc, verify_k - Kc
else:
n_cond = ((verify_k + G - 1) // G) * G # round up to group
n_par = 0
self.reset_cache()
h_seed = self.prefill(prompt_ids) # [1,d] hidden at pos-1
pos = len(prompt_ids)
committed = list(prompt_ids)
out = []
pending = None # committed, not yet forwarded
n_rounds = 0
while len(out) < new_tokens:
first = 0 if pending is None else 1
# group-0 conditioning = last G tokens ending AT the anchor
# (anchor = last forwarded pos; pending sits one past it)
if pending is None:
hist = committed[-G:]
else:
hist = committed[-G - 1:-1]
if len(hist) < G:
hist = [committed[0]] * (G - len(hist)) + hist
hist_t = torch.tensor([hist], device=self.device)
draft = self.model.spec_draft(h_seed, hist_t, first_head=first,
pending_id=pending,
n_cond=n_cond, n_par=n_par)
inp = ([pending] if pending is not None else []) + draft[0].tolist()
if verify_k is not None:
inp = inp[:verify_k]
L = len(inp)
ids = torch.tensor([inp], device=self.device, dtype=torch.long)
# verify: causal forward over the draft chunk
h_all = self.model(ids, self._cache_k, self._cache_v, pos,
causal_extend=True) # [1, L, d]
# base argmax for each input position's next token:
# pos-1 -> h_seed, pos+i -> h_all[i-1]
prev_h = torch.cat([h_seed.unsqueeze(1), h_all[:, :-1]], dim=1)
preds = self.model.lm_head(prev_h).argmax(-1)[0] # [L]
mism = (ids[0] != preds).nonzero()
j = int(mism[0].item()) if mism.numel() else L # first mismatch
corr = int((preds[j] if j < L
else self.model.lm_head(h_all[:, -1]).argmax(-1)).item())
out.extend(inp[first:j])
out.append(corr)
committed.extend(inp[first:j])
committed.append(corr)
if j > 0:
h_seed = h_all[:, j - 1] # hidden at pos+j-1
pending = corr
pos += j
n_rounds += 1
stats = {"rounds": n_rounds,
"avg_commit": len(out) / max(1, n_rounds)}
return out[:new_tokens], stats
# --------------------------------------------------------------------------------------
# FastEngine: ALiBi + sliding window + CUDA graphs
# --------------------------------------------------------------------------------------
class FastEngine:
"""Inference engine with fixed sliding-window cache + CUDA graph capture.
- Fixed-size KV cache [1, W, n_kv, hd] -> attention is always [K+1, W], never grows
- ALiBi position bias (precomputed, fixed) -> static shapes for CUDA graphs
- CUDA graph captures the entire forward pass -> 1 kernel launch per step
- Falls back to eager if graph capture fails
The prefill (prompt processing) runs eagerly. Only the generation loop is graphed.
"""
def __init__(self, model: SpecModel, device="cuda"):
self.model = model
self.cfg = model.cfg
self.device = device
self.dtype = model.tok_emb.weight.dtype
self.K1 = self.cfg.medusa_heads + 1
self.W = self.cfg.window_size if self.cfg.window_size > 0 else self.cfg.max_seq_len
# Precompute medusa argmax/argmin for r=1 fast path
self._medusa_fast = False
if self.cfg.medusa_rank == 1:
model.precompute_medusa_tokens()
self._medusa_fast = True
# Fixed-size cache (window W, not max_seq_len)
self._cache_k = [
torch.zeros(1, self.W, self.cfg.n_kv_heads, self.cfg.head_dim,
device=self.device, dtype=self.dtype)
for _ in range(self.cfg.n_layers)
]
self._cache_v = [t.clone() for t in self._cache_k]
# CUDA graph state
self._graph = None
self._static_input = None
self._static_output = None
def reset_cache(self):
for t in self._cache_k:
t.zero_()
for t in self._cache_v:
t.zero_()
@torch.no_grad()
def capture_graph(self):
"""Capture the K+1-token forward pass as a CUDA graph. Falls back to eager."""
try:
# static input tensor [1, K+1]
self._static_input = torch.zeros(1, self.K1, device=self.device, dtype=torch.long)
# warmup (3 iters to initialize lazy state, cuDNN, etc.)
for _ in range(3):
_ = self.model(self._static_input, self._cache_k, self._cache_v,
sliding_window=True)
torch.cuda.synchronize()
# capture
self._graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(self._graph):
self._static_output = self.model(self._static_input, self._cache_k,
self._cache_v, sliding_window=True)
torch.cuda.synchronize()
return True
except Exception as e:
print(f" CUDA graph capture failed ({type(e).__name__}: {e}), using eager")
self._graph = None
return False
@torch.no_grad()
def step(self, input_ids: torch.Tensor) -> torch.Tensor:
"""One forward pass over [1, K+1] tokens. Returns hidden [1, K+1, d]."""
if self._graph is not None:
self._static_input.copy_(input_ids)
self._graph.replay()
return self._static_output
return self.model(input_ids, self._cache_k, self._cache_v, sliding_window=True)
@torch.no_grad()
def next_tokens_from_hidden(self, h_last: torch.Tensor):
main = self.model.lm_head(h_last).argmax(-1)
if self._medusa_fast:
medusa = self.model.medusa_argmax_fast(h_last)
else:
medusa = self.model.medusa_argmax(h_last)
return main, medusa
@torch.no_grad()
def prefill(self, prompt_ids: list[int]) -> torch.Tensor:
"""Process prompt eagerly (no graph). Writes K,V at position 0 in the window."""
P = len(prompt_ids)
if P > self.W:
prompt_ids = prompt_ids[-self.W:] # only keep last W tokens
P = self.W
ids = torch.tensor([prompt_ids], device=self.device, dtype=torch.long)
# Use non-sliding-window path: writes at start_pos=0, standard causal attention
h = self.model(ids, self._cache_k, self._cache_v, start_pos=0,
sliding_window=False)
return h[:, -1]
@torch.no_grad()
def speculative(self, prompt_ids: list[int], new_tokens: int) -> list[int]:
"""K+1 tokens per forward pass, all accepted. Uses CUDA graph if captured."""
self.reset_cache()
K = self.cfg.medusa_heads
h_last = self.prefill(prompt_ids)
out = []
main, medusa = self.next_tokens_from_hidden(h_last)
accepted = torch.cat([main, medusa[0]])
out.extend(accepted.tolist())
while len(out) < new_tokens:
h = self.step(accepted.unsqueeze(0))
h_last = h[:, -1]
main, medusa = self.next_tokens_from_hidden(h_last)
accepted = torch.cat([main, medusa[0]])
out.extend(accepted.tolist())
return out[:new_tokens]
# --------------------------------------------------------------------------------------
# benchmarking helper
# --------------------------------------------------------------------------------------
def benchmark(engine: InferenceEngine, prompt_ids: list[int], new_tokens: int,
mode: str, warmup=3, repeats=5):
"""Returns (tokens, median tokens/sec)."""
fn = engine.autoregressive if mode == "ar" else engine.speculative
# warmup
for _ in range(warmup):
fn(prompt_ids, new_tokens)
torch.cuda.synchronize()
times = []
for _ in range(repeats):
torch.cuda.synchronize()
t0 = time.perf_counter()
toks = fn(prompt_ids, new_tokens)
torch.cuda.synchronize()
times.append(time.perf_counter() - t0)
times.sort()
med = times[len(times) // 2]
return toks, new_tokens / med
# --------------------------------------------------------------------------------------
# v5 Engine: compressed MoBA + FP8 cache + real chained speculative prefill
# --------------------------------------------------------------------------------------
class V5Engine:
"""Inference engine for the v5 architecture.
Features:
- Compressed KV cache (DeepSeek V4-style): every m tokens -> 1 KV entry
- FP8 KV cache storage (E4M3) for 2x memory reduction
- MoBA block-sparse attention: top-k block selection (fused SDPA)
- Within-block sparse attention: attend to every Nth compressed entry
- Real chained speculative prefill: each step feeds Medusa output back as
input, writes actual K,V to cache (not pre-filled random data)
- N-gram draft extension: free tokens per step via lookup table
- torch.compile: fused kernels for FFN + attention
The prefill processes the prompt in K+1-token chunks. Each chunk:
1. Forward pass through compressed MoBA -> hidden states
2. Medusa heads predict K future tokens from last hidden state
3. N-gram extends with M more free tokens (lookup, no forward pass)
4. Those K+M+1 tokens become the NEXT chunk's input
5. Actual K,V from step 1 are compressed and written to cache
"""
def __init__(self, model: SpecModel, device="cuda", ngram=None):
self.model = model
self.cfg = model.cfg
self.device = device
self.dtype = model.tok_emb.weight.dtype
self.K = self.cfg.medusa_heads
self.K1 = self.K + 1
self.m = self.cfg.kv_compress_m
self.use_fp8 = self.cfg.kv_fp8
self.BS_comp = self.cfg.block_size_comp # compressed entries per block
self.cache_dt = (torch.float8_e4m3fn if self.use_fp8
else self.dtype)
self.ngram = ngram
self.ngram_extend = 0 # n-gram tokens to add per step (0 = disabled)
# Precompute medusa argmax/argmin for r=1 fast path
self._medusa_fast = False
if self.cfg.medusa_rank == 1:
model.precompute_medusa_tokens()
self._medusa_fast = True
# Block-structured compressed KV cache (allocated on first use)
self._cache_k = None # per-layer compressed K blocks
self._cache_v = None # per-layer compressed V blocks
self._k_bar_comp = None # per-layer block-mean K (gating)
self._ring_k = None # per-layer raw-K ring [1, raw_window, n_kv, hd]
self._ring_v = None
self._ring_len = 0
self._n_blocks = 0
self._n_filled_blocks = 0
self._current_block_fill = 0
# torch.compile state
self._compiled = False
self._compiled_step = None
# CUDA-graph state for spec_draft (v6 heads): the 32-group sequential
# draft is launch-bound in eager (~0.6s); graphed it replays in ~ms.
self._draft_graph = None
self._draft_h = None
self._draft_hist = None
self._draft_out = None
def alloc_cache(self, n_blocks: int):
"""Allocate block-structured compressed KV cache.
Each block holds BS_comp compressed entries.
Cache shape: [1, n_blocks, BS_comp, n_kv, hd] per layer.
"""
cfg = self.cfg
self._n_blocks = n_blocks
shape = (1, n_blocks, self.BS_comp, cfg.n_kv_heads, cfg.head_dim)
self._cache_k = [torch.zeros(shape, device=self.device,
dtype=self.cache_dt)
for _ in range(cfg.n_layers)]
self._cache_v = [torch.zeros(shape, device=self.device,
dtype=self.cache_dt)
for _ in range(cfg.n_layers)]
self._k_bar_comp = [
torch.zeros(1, n_blocks, cfg.n_kv_heads, cfg.head_dim,
device=self.device, dtype=self.dtype)
for _ in range(cfg.n_layers)
]
W = cfg.raw_window
self._ring_k = [
torch.zeros(1, W, cfg.n_kv_heads, cfg.head_dim,
device=self.device, dtype=self.dtype)
for _ in range(cfg.n_layers)
]
self._ring_v = [
torch.zeros(1, W, cfg.n_kv_heads, cfg.head_dim,
device=self.device, dtype=self.dtype)
for _ in range(cfg.n_layers)
]
self._ring_len = 0
def reset_cache(self):
if self._cache_k is not None:
for t in self._cache_k + self._cache_v + self._k_bar_comp:
t.zero_()
self._n_filled_blocks = 0
self._current_block_fill = 0
self._ring_len = 0
@torch.no_grad()
def try_compile(self):
"""Compile just the FFN compute (65% of step time, fully static shapes).
Uses a standalone function with raw weight tensors to avoid the
nn.Module child replacement issue.
"""
if self._compiled:
return True
import torch.nn.functional as F
# Compile a standalone SwiGLU function that takes raw tensors
d_model = self.cfg.d_model
@torch.compile(mode="max-autotune", dynamic=False)
def _compiled_swiglu(x, w_gate, b_gate, w_up, b_up, w_down, b_down, norm_w):
h = F.rms_norm(x, (d_model,), norm_w)
gate = F.linear(h, w_gate, b_gate)
up = F.linear(h, w_up, b_up)
return F.linear(F.silu(gate) * up, w_down, b_down)
# Attach compiled FFN to each block
for blk in self.model.blocks:
ffn = blk.ffn
norm = blk.norm2
wg, wu, wd = ffn.w_gate, ffn.w_up, ffn.w_down
bg = ffn.w_gate.bias if ffn.w_gate.bias is not None else None
bu = ffn.w_up.bias if ffn.w_up.bias is not None else None
bd = ffn.w_down.bias if ffn.w_down.bias is not None else None
def make_fn(wg, bg, wu, bu, wd, bd, nw):
def fn(x):
return _compiled_swiglu(x, wg, bg, wu, bu, wd, bd, nw)
return fn
blk._compiled_ffn_fn = make_fn(wg.weight, bg, wu.weight, bu,
wd.weight, bd, norm.weight)
self._compiled = True
return True
@torch.no_grad()
def _step(self, input_ids: torch.Tensor) -> torch.Tensor:
"""One v5 forward pass. Writes compressed K,V to current block.
Returns hidden states [1, T, d]."""
h = self.model.forward_moba_v5(
input_ids, self._cache_k, self._cache_v, self._k_bar_comp,
self._n_filled_blocks, self._current_block_fill,
self._ring_k, self._ring_v, self._ring_len)
T = input_ids.shape[1]
self._ring_len = min(self.cfg.raw_window, self._ring_len + T)
# Advance block position (model handles overflow internally)
n_new = T // self.m # compressed entries written
self._current_block_fill += n_new
while self._current_block_fill >= self.BS_comp:
self._n_filled_blocks += 1
self._current_block_fill -= self.BS_comp
return h
@torch.no_grad()
def _next_tokens(self, h_last: torch.Tensor, hist_ids: torch.Tensor = None):
"""Given hidden at last position, return (main, draft) tokens.
r=1: medusa_argmax[_fast]. r>1 (v6): spec_draft conditioned chain."""
main = self.model.lm_head(h_last).argmax(-1) # [1]
if self.cfg.medusa_rank == 1:
if self._medusa_fast:
medusa = self.model.medusa_argmax_fast(h_last) # [1, K]
else:
medusa = self.model.medusa_argmax(h_last) # [1, K]
return main, medusa
# v6 hybrid heads: conditioned chain + parallel tail
G = self.cfg.medusa_cond_group
if hist_ids is None:
hist_ids = torch.zeros(1, G, dtype=torch.long, device=self.device)
draft = self._graphed_draft(h_last, hist_ids) # [1,K]
return main, draft
@torch.no_grad()
def _graphed_draft(self, h_last: torch.Tensor, hist_ids: torch.Tensor):
"""spec_draft via CUDA graph (falls back to eager on failure)."""
if self._draft_graph is None:
try:
self._draft_h = torch.zeros_like(h_last)
self._draft_hist = torch.zeros_like(hist_ids)
self._draft_h.copy_(h_last)
self._draft_hist.copy_(hist_ids)
for _ in range(3): # warmup
self.model.spec_draft(self._draft_h, self._draft_hist,
first_head=0)
torch.cuda.synchronize()
g = torch.cuda.CUDAGraph()
with torch.cuda.graph(g):
self._draft_out = self.model.spec_draft(
self._draft_h, self._draft_hist, first_head=0)
torch.cuda.synchronize()
self._draft_graph = g
except Exception as e:
print(f" draft graph capture failed ({type(e).__name__}: {e})"
f" — eager fallback")
self._draft_graph = False
return self.model.spec_draft(h_last, hist_ids, first_head=0)
if self._draft_graph:
self._draft_h.copy_(h_last)
self._draft_hist.copy_(hist_ids)
self._draft_graph.replay()
return self._draft_out
return self.model.spec_draft(h_last, hist_ids, first_head=0)
def _ngram_extend(self, context_tokens: list[int], n_tokens: int) -> list[int]:
"""Use n-gram to predict n_tokens for free (no forward pass).
Returns predicted tokens (may be shorter if n-gram runs out)."""
if self.ngram is None or self.ngram_extend == 0:
return []
return self.ngram.predict_batch(context_tokens, n_tokens)
@torch.no_grad()
def prefill_chained(self, prompt_ids: list[int], max_tokens: int = None,
prefill_chunk: int = None) -> dict:
"""REAL chained speculative prefill with n-gram extension.
Processes the prompt in chunks. While real prompt tokens remain, chunks
are `prefill_chunk` tokens (default K+1; larger = faster — the draft is
skipped so chunk size is free). Once the prompt is exhausted, each step
chains K+1 drafted tokens (+ optional n-gram extension).
This is a REAL prefill: the cache is built from actual model K,V.
Returns dict with timing stats.
"""
cfg = self.cfg
K1 = self.K1
m = self.m
ngram_m = self.ngram_extend
pc = prefill_chunk or K1
pc -= pc % m # pooling needs T divisible by m
# First chunk of the prompt
ids = torch.tensor([prompt_ids[:pc]], device=self.device, dtype=torch.long)
remaining_prompt = prompt_ids[pc:]
total_processed = len(ids[0])
total_ngram_tokens = 0
# Allocate cache: estimate blocks needed (only if not already allocated)
if max_tokens is None:
max_tokens = len(prompt_ids)
n_blocks_needed = (max_tokens // self.m // self.BS_comp) + 2
if self._cache_k is None or self._n_blocks < n_blocks_needed:
self.alloc_cache(n_blocks_needed)
self.reset_cache()
t0 = time.perf_counter()
n_steps = 0
# Process prompt in chunks
hist_buf = list(prompt_ids[-cfg.medusa_cond_group:]) if cfg.medusa_rank > 1 else None
while total_processed < len(prompt_ids):
h = self._step(ids) # writes compressed K,V, advances block pos
if len(remaining_prompt) >= pc:
# Pure prefill: real tokens available, skip the draft entirely
# (v6 spec_draft is ~0.6s eager — running it here is pure waste)
chunk = remaining_prompt[:pc]
remaining_prompt = remaining_prompt[pc:]
if hist_buf is not None:
hist_buf.extend(chunk)
ids = torch.tensor([chunk], device=self.device, dtype=torch.long)
else:
h_last = h[:, -1] # [1, d]
hist_t = None
if hist_buf is not None:
hist_t = torch.tensor([hist_buf[-cfg.medusa_cond_group:]],
device=self.device)
main, medusa = self._next_tokens(h_last, hist_t)
# Next chunk: use remaining prompt if available, else Medusa + n-gram
next_tokens = torch.cat([main, medusa[0]]) # [K+1]
if remaining_prompt:
# Tail of prompt (< pc tokens): real tokens + draft padding
chunk = remaining_prompt[:pc]
remaining_prompt = remaining_prompt[pc:]
if len(chunk) < pc:
chunk = chunk + next_tokens[len(chunk):].tolist()
if hist_buf is not None:
hist_buf.extend(chunk)
ids = torch.tensor([chunk], device=self.device, dtype=torch.long)
else:
# No more prompt: use Medusa + n-gram extension
ids = next_tokens.unsqueeze(0)
if hist_buf is not None:
hist_buf.extend(next_tokens.tolist())
# N-gram extension: free tokens, no forward pass needed
if ngram_m > 0:
ctx = ids[0].tolist()
ngram_tokens = self._ngram_extend(ctx, ngram_m)
total_ngram_tokens += len(ngram_tokens)
total_processed += ids.shape[1]
n_steps += 1
# Final step for the last chunk
h = self._step(ids)
n_steps += 1
total_processed += ids.shape[1]
elapsed = time.perf_counter() - t0
return {
"tokens": total_processed,
"ngram_tokens": total_ngram_tokens,
"steps": n_steps,
"time": elapsed,
"tok/s": total_processed / elapsed if elapsed > 0 else 0,
"n_filled_blocks": self._n_filled_blocks,
"current_block_fill": self._current_block_fill,
}
@torch.no_grad()
def generate(self, prompt_ids: list[int], new_tokens: int) -> list[int]:
"""Generate new_tokens after prefill. Uses compressed MoBA cache."""
# Allocate cache for prefill + generation
total_tokens = len(prompt_ids) + new_tokens
n_blocks_needed = (total_tokens // self.m // self.BS_comp) + 4
self.alloc_cache(n_blocks_needed)
self.reset_cache()
# Prefill first (big chunks — draft is skipped while real tokens remain)
self.prefill_chained(prompt_ids, prefill_chunk=4096)
# Get last hidden state for first draft
ids = torch.tensor([[prompt_ids[-1]]], device=self.device, dtype=torch.long)
out_parts = []
n_out = 0
G = self.cfg.medusa_cond_group
hist_t = (torch.tensor([prompt_ids[-G:]], device=self.device)
if self.cfg.medusa_rank > 1 else None)
while n_out < new_tokens:
h = self._step(ids)
h_last = h[:, -1]
main, medusa = self._next_tokens(h_last, hist_t)
accepted = torch.cat([main, medusa[0]]) # [K+1]
out_parts.append(accepted)
n_out += accepted.numel()
if self.cfg.medusa_rank > 1:
hist_t = accepted[-G:].unsqueeze(0) # stays on GPU
ids = accepted.unsqueeze(0)
return torch.cat(out_parts)[:new_tokens].tolist()
@torch.no_grad()
def generate_with_ngram(self, prompt_ids: list[int], new_tokens: int,
ngram_extend: int = 4096) -> dict:
"""Generate with n-gram draft extension (FREE tokens, no forward pass).
Each step:
1. Forward pass: K+1 tokens (main + medusa) — costs ~20ms
2. N-gram extension: M tokens from lookup table — costs ~0.5ms
3. Total: K+1+M tokens per step
The n-gram tokens are NOT fed back into the model (accept-all mode).
This means the model generates K+1 tokens per step, and the n-gram
adds M free tokens on top. Effective throughput: (K+1+M) / step_time.
Returns dict with output tokens and timing stats.
"""
# Allocate cache for prefill + generation
total_tokens = len(prompt_ids) + new_tokens
n_blocks_needed = (total_tokens // self.m // self.BS_comp) + 4
self.alloc_cache(n_blocks_needed)
self.reset_cache()
# Prefill first (big chunks — draft is skipped while real tokens remain)
self.prefill_chained(prompt_ids, prefill_chunk=4096)
ids = torch.tensor([[prompt_ids[-1]]], device=self.device, dtype=torch.long)
out = []
hist_buf = list(prompt_ids)
ngram_tokens_total = 0
n_steps = 0
t0 = time.perf_counter()
while len(out) < new_tokens:
h = self._step(ids)
h_last = h[:, -1]
hist_t = None
if self.cfg.medusa_rank > 1:
hist_t = torch.tensor(
[hist_buf[-self.cfg.medusa_cond_group:]],
device=self.device)
main, medusa = self._next_tokens(h_last, hist_t)
accepted = torch.cat([main, medusa[0]]) # [K+1]
model_tokens = accepted.tolist()
out.extend(model_tokens)
hist_buf.extend(model_tokens)
# N-gram extension: predict M free tokens from recent context
if self.ngram is not None and ngram_extend > 0:
# Use the full output history as context (includes prompt tokens
# at the start, which are real text). For untrained models,
# the model output is garbage, so n-gram predictions will be
# limited. After training, model output will look like real
# text and n-gram predictions will be much longer.
ctx = out[-self.ngram.n:] if len(out) >= self.ngram.n else \
(list(prompt_ids[-(self.ngram.n - len(out)):]) + out)
ngram_tokens = self.ngram.predict_batch(ctx, ngram_extend)
out.extend(ngram_tokens)
ngram_tokens_total += len(ngram_tokens)
# Next step input: only the model's K+1 tokens (not n-gram)
ids = accepted.unsqueeze(0)
n_steps += 1
elapsed = time.perf_counter() - t0
return {
"tokens": len(out[:new_tokens]),
"model_tokens": n_steps * self.K1,
"ngram_tokens": ngram_tokens_total,
"steps": n_steps,
"time": elapsed,
"tok/s": len(out[:new_tokens]) / elapsed if elapsed > 0 else 0,
"output": out[:new_tokens],
}
# ---------------------------------------------------- verified @ long ctx
@torch.no_grad()
def speculative_verified(self, prompt_ids: list[int], new_tokens: int,
verify_k: int = None, conf_tau: float = 0.0,
prefill_chunk: int = 4096):
"""Greedy-verified speculative decoding on the compressed-MoBA cache.
Same commit rule as InferenceEngine.speculative_verified — output is
the base model's greedy decode under the MoBA-approx attention path —
but the compressed cache only rewinds at m-token entry granularity.
Committed tokens that don't fill a whole compressed entry form a
`queue` carried into the next round (generalizes the single pending
token), so no committed token is ever dropped or double-counted.
verify_k caps the draft length per round (cheaper rounds when the
accept streak is short).
"""
cfg = self.cfg
assert cfg.medusa_rank > 1, "verified mode needs v6 heads"
G, m, BS = cfg.medusa_cond_group, cfg.kv_compress_m, self.BS_comp
Kc, Kp = cfg.medusa_cond_heads, cfg.medusa_par_heads
if verify_k is None or verify_k >= Kc + Kp:
n_cond, n_par = Kc, Kp
elif verify_k >= Kc:
n_cond, n_par = Kc, verify_k - Kc
else:
n_cond = ((verify_k + G - 1) // G) * G
n_par = 0
# --- chunked prefill; capture h at the last cache-covered position ---
total = len(prompt_ids) + new_tokens + Kc + Kp + 8
n_blocks_needed = (total // m // BS) + 4
self.alloc_cache(n_blocks_needed)
self.reset_cache()
committed = list(prompt_ids)
pc = prefill_chunk - prefill_chunk % m
h_seed = None
i = 0
while i < len(prompt_ids):
chunk = prompt_ids[i:i + pc]
ids = torch.tensor([chunk], device=self.device, dtype=torch.long)
h = self.model.forward_moba_v5(
ids, self._cache_k, self._cache_v, self._k_bar_comp,
self._n_filled_blocks, self._current_block_fill,
self._ring_k, self._ring_v, self._ring_len)
n_new = len(chunk) // m
if n_new > 0:
h_seed = h[:, n_new * m - 1] # hidden at last covered
self._ring_len = min(cfg.raw_window, self._ring_len + len(chunk))
self._current_block_fill += n_new
while self._current_block_fill >= BS:
self._n_filled_blocks += 1
self._current_block_fill -= BS
i += len(chunk)
fwd = (self._n_filled_blocks * BS + self._current_block_fill) * m
queue = committed[fwd:] # committed, not in cache
if h_seed is None: # tiny prompt: forward it
ids = torch.tensor([prompt_ids], device=self.device, dtype=torch.long)
h_seed = self.model.forward_moba_v5(
ids, self._cache_k, self._cache_v, self._k_bar_comp,
self._n_filled_blocks, self._current_block_fill,
self._ring_k, self._ring_v, self._ring_len)[:, -1]
self._ring_len = min(cfg.raw_window,
self._ring_len + len(prompt_ids))
out = []
n_rounds = 0
while len(out) < new_tokens:
Q = len(queue)
hist = committed[max(0, fwd - G):fwd]
hist = [committed[0]] * (G - len(hist)) + hist
hist_t = torch.tensor([hist], device=self.device)
prefix_t = (torch.tensor([queue], device=self.device)
if Q else None)
draft, conf = self.model.spec_draft(
h_seed, hist_t, first_head=Q, prefix_ids=prefix_t,
n_cond=n_cond, n_par=n_par, return_conf=True)
vd = n_cond + n_par - Q
if conf_tau > 0:
# dspark scheduled verification: only verify while the
# cumulative acceptance probability stays above tau —
# don't waste a verify forward on a doomed suffix.
surv = conf[0].cumprod(0)
ok = (surv > conf_tau).nonzero()
vlen = int(ok[-1].item()) + 1 if ok.numel() else 1
vd = min(vd, vlen)
inp = queue + draft[0].tolist()[:vd]
L = len(inp)
ids = torch.tensor([inp], device=self.device, dtype=torch.long)
snap_blocks, snap_fill = (self._n_filled_blocks,
self._current_block_fill)
snap_ring = self._ring_len
h_all = self.model.forward_moba_v5(
ids, self._cache_k, self._cache_v, self._k_bar_comp,
snap_blocks, snap_fill,
self._ring_k, self._ring_v, snap_ring) # [1, L, d]
# verify: base argmax at each position = h_seed then h_all[:-1]
prev_h = torch.cat([h_seed.unsqueeze(1), h_all[:, :-1]], dim=1)
preds = self.model.lm_head(prev_h).argmax(-1)[0] # [L]
dt_ = torch.tensor(inp[Q:], device=self.device)
mism = (dt_ != preds[Q:]).nonzero()
jd = int(mism[0].item()) if mism.numel() else len(dt_)
corr = int((preds[Q + jd] if Q + jd < L
else self.model.lm_head(h_all[:, -1])
.argmax(-1)).item())
committed.extend(inp[Q:Q + jd])
committed.append(corr)
out.extend(inp[Q:Q + jd])
out.append(corr)
# rewind cache to last fully-committed m-entry boundary.
# NOTE: corr's cache position held the REJECTED draft token during
# this forward, so its entry must NOT be kept — corr stays queued.
keep_e = min((Q + jd) // m, L // m)
E_next = snap_blocks * BS + snap_fill + keep_e
self._n_filled_blocks = E_next // BS
self._current_block_fill = E_next % BS
fwd = E_next * m
queue = committed[fwd:]
if keep_e > 0:
h_seed = h_all[:, keep_e * m - 1] # hidden at fwd-1
# same rewind for the raw ring: keep committed tokens only
dropped = max(0, snap_ring + L - cfg.raw_window)
self._ring_len = min(cfg.raw_window,
snap_ring - dropped + Q + jd)
n_rounds += 1
stats = {"rounds": n_rounds,
"avg_commit": len(out) / max(1, n_rounds)}
return out[:new_tokens], stats