Download code/inference.py from Akahsizrr/spec100m: direct link, hf CLI and curl.
- Browser
- Download file 42 kB
-
https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/inference.py
- Command line
-
hf download hf://Akahsizrr/spec100m/code/inference.py
-
curl -L -o inference.py https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/inference.py
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 | |
| 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 | |
| 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) | |
| 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 | |
| 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 | |
| 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) | |
| 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) | |
| 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_() | |
| 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 | |
| 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) | |
| 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 | |
| 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] | |
| 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 | |
| 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 | |
| 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 | |
| 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 | |
| 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 | |
| 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) | |
| 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, | |
| } | |
| 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() | |
| 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 | |
| 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 | |