"""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