Download code/model.py from Akahsizrr/spec100m: direct link, hf CLI and curl.
- Browser
- Download file 53.6 kB
-
https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/model.py
- Command line
-
hf download hf://Akahsizrr/spec100m/code/model.py
-
curl -L -o model.py https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/model.py
53.6 kB
| """Transformer + Medusa speculative heads. | |
| - RMSNorm, RoPE, GQA, SwiGLU, tied embeddings. | |
| - SDPA (torch built-in fused attention) -> no flash-attn install pain on Windows. | |
| - Supports training (full-sequence, is_causal) and inference (KV-cache extend). | |
| - forward() returns last-layer hidden states; LM head + Medusa heads are separate | |
| so the speculative decoder can call them once on the final hidden state. | |
| """ | |
| from __future__ import annotations | |
| import math | |
| from dataclasses import dataclass | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from config import Config | |
| # enable_gqa added in torch 2.5 — detect via version (signature() fails on C builtin) | |
| _tv = torch.__version__.split("+")[0].split(".") | |
| _HAS_ENABLE_GQA = (int(_tv[0]), int(_tv[1])) >= (2, 5) | |
| def _gqa_kw(): | |
| return {"enable_gqa": True} if _HAS_ENABLE_GQA else {} | |
| def _expand_kv(q, k, v): | |
| """On torch <2.5 (no enable_gqa), expand K/V head dim to match Q.""" | |
| if _HAS_ENABLE_GQA or k.shape[1] == q.shape[1]: | |
| return k, v | |
| B, n_kv, S, hd = k.shape | |
| rep = q.shape[1] // n_kv | |
| k = k[:, :, None].expand(B, n_kv, rep, S, hd).reshape(B, q.shape[1], S, hd) | |
| v = v[:, :, None].expand(B, n_kv, rep, S, hd).reshape(B, q.shape[1], S, hd) | |
| return k, v | |
| # -------------------------------------------------------------------------------------- | |
| # RoPE | |
| # -------------------------------------------------------------------------------------- | |
| def precompute_rope(head_dim: int, max_seq: int, base: float, device, dtype) -> torch.Tensor: | |
| """Returns cos, sin of shape [max_seq, head_dim/2] each (real).""" | |
| half = head_dim // 2 | |
| freqs = 1.0 / (base ** (torch.arange(0, half, device=device, dtype=torch.float32) / half)) | |
| t = torch.arange(max_seq, device=device, dtype=torch.float32) | |
| ang = torch.outer(t, freqs) # [max_seq, half] | |
| return ang.cos().to(dtype), ang.sin().to(dtype) | |
| def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: | |
| """x: [B, T, H, D]. cos/sin: [T, D/2]. Returns rotated x.""" | |
| B, T, H, D = x.shape | |
| x1, x2 = x[..., :D // 2], x[..., D // 2:] # split into even/odd pairs (interleaved) | |
| # cos/sin broadcast over [T,1,D/2] -> [1,T,1,D/2] | |
| c = cos[:T].view(1, T, 1, D // 2) | |
| s = sin[:T].view(1, T, 1, D // 2) | |
| rot = torch.cat([x1 * c - x2 * s, x1 * s + x2 * c], dim=-1) | |
| return rot.to(x.dtype) | |
| # -------------------------------------------------------------------------------------- | |
| # ALiBi (fixed additive bias, no rotation — for sliding window + CUDA graphs) | |
| # -------------------------------------------------------------------------------------- | |
| def precompute_alibi_bias(n_heads: int, q_len: int, kv_len: int, dtype, device) -> torch.Tensor: | |
| """Fixed ALiBi attention bias for a sliding window. | |
| Queries are the last q_len positions of a kv_len-size window. | |
| Returns [n_heads, q_len, kv_len] additive bias (0 or negative). | |
| bias[h, i, j] = -slope_h * (q_pos_i - j) if j <= q_pos_i (can attend) | |
| -inf if j > q_pos_i (causal mask) | |
| where q_pos_i = kv_len - q_len + i (query i is at window position kv_len-q_len+i). | |
| """ | |
| # ALiBi slopes: geometric sequence per head | |
| slopes = 1.0 / (2.0 ** (torch.arange(1, n_heads + 1, device=device, dtype=torch.float32) / n_heads)) | |
| slopes = slopes.view(n_heads, 1, 1) # [H, 1, 1] | |
| q_pos = torch.arange(kv_len - q_len, kv_len, device=device, dtype=torch.float32) # [q_len] | |
| k_pos = torch.arange(kv_len, device=device, dtype=torch.float32) # [kv_len] | |
| dist = q_pos.unsqueeze(1) - k_pos.unsqueeze(0) # [q_len, kv_len] | |
| bias = -slopes * dist.clamp(min=0) # [H, q_len, kv_len], 0 for future positions | |
| causal = torch.where(dist.unsqueeze(0) >= 0, bias, torch.tensor(float("-inf"), device=device)) | |
| return causal.to(dtype) | |
| # -------------------------------------------------------------------------------------- | |
| # Attention with GQA + KV cache | |
| # -------------------------------------------------------------------------------------- | |
| class Attention(nn.Module): | |
| def __init__(self, cfg: Config): | |
| super().__init__() | |
| self.cfg = cfg | |
| self.n_heads = cfg.n_heads | |
| self.n_kv = cfg.n_kv_heads | |
| self.hd = cfg.head_dim | |
| d = cfg.d_model | |
| self.wq = nn.Linear(d, self.n_heads * self.hd, bias=False) | |
| self.wk = nn.Linear(d, self.n_kv * self.hd, bias=False) | |
| self.wv = nn.Linear(d, self.n_kv * self.hd, bias=False) | |
| self.wo = nn.Linear(self.n_heads * self.hd, d, bias=False) | |
| self.scale = 1.0 / math.sqrt(self.hd) | |
| # MoBA params | |
| self.block_size = cfg.moba_block_size | |
| self.top_k = cfg.moba_top_k | |
| # v5: KV compression params (DeepSeek V4-style) | |
| self.compress_m = cfg.kv_compress_m | |
| self.use_fp8 = cfg.kv_fp8 | |
| self.sparse_stride = cfg.sparse_stride | |
| # Compression projections: hidden -> compressed KV entry (shared K=V, MQA-style) | |
| # W_KV: d -> n_kv*hd, W_Z: d -> n_kv*hd (compression weights), B: m x n_kv*hd (positional bias) | |
| if self.compress_m > 1: | |
| self.w_kv_comp = nn.Linear(d, self.n_kv * self.hd, bias=False) | |
| self.w_v_comp = nn.Linear(d, self.n_kv * self.hd, bias=False) | |
| self.w_z_comp = nn.Linear(d, self.n_kv * self.hd, bias=False) | |
| self.comp_bias = nn.Parameter(torch.zeros(self.compress_m, self.n_kv * self.hd)) | |
| def forward(self, x, cos, sin, cache_k=None, cache_v=None, start_pos=0, | |
| alibi_bias=None, sliding_window=False, causal_extend=False): | |
| """ | |
| x: [B, T, d] | |
| cache_k/v: [B, S, n_kv, hd] preallocated. | |
| start_pos: int, current cached length (growing cache) or ignored (sliding window). | |
| alibi_bias: [n_heads, T, W] precomputed ALiBi bias (None = use RoPE). | |
| sliding_window: if True, cache is fixed-size W; roll by T and write at end. | |
| causal_extend: if True, use a real causal mask within the new chunk | |
| (needed for speculative verification — accept-all path skips it | |
| for speed). | |
| Returns: [B, T, d] | |
| """ | |
| B, T, _ = x.shape | |
| q = self.wq(x).view(B, T, self.n_heads, self.hd) | |
| k = self.wk(x).view(B, T, self.n_kv, self.hd) | |
| v = self.wv(x).view(B, T, self.n_kv, self.hd) | |
| if alibi_bias is None: | |
| if cos is not None and cos.numel() > 0: | |
| off = 0 if sliding_window else start_pos | |
| q = apply_rope(q, cos[off:off + T], sin[off:off + T]) | |
| k = apply_rope(k, cos[off:off + T], sin[off:off + T]) | |
| # else: no position encoding (prefill with ALiBi config — quality irrelevant) | |
| if cache_k is not None: | |
| if sliding_window: | |
| # Fixed-size window: roll left by T, write new K/V at the end. | |
| # cache_k is [B, W, n_kv, hd]; always attend to full W. | |
| cache_k.copy_(torch.roll(cache_k, shifts=-T, dims=1)) | |
| cache_v.copy_(torch.roll(cache_v, shifts=-T, dims=1)) | |
| cache_k[:, -T:] = k | |
| cache_v[:, -T:] = v | |
| k = cache_k | |
| v = cache_v | |
| S = cache_k.shape[1] | |
| else: | |
| cache_k[:, start_pos:start_pos + T] = k | |
| cache_v[:, start_pos:start_pos + T] = v | |
| k = cache_k[:, :start_pos + T] | |
| v = cache_v[:, :start_pos + T] | |
| S = start_pos + T | |
| else: | |
| S = T | |
| # reshape for SDPA: [B, H, T, hd] | |
| q = q.transpose(1, 2) # [B, n_heads, T, hd] | |
| k = k.transpose(1, 2) # [B, n_kv, S, hd] | |
| v = v.transpose(1, 2) # [B, n_kv, S, hd] | |
| k, v = _expand_kv(q, k, v) | |
| if alibi_bias is not None: | |
| # ALiBi: use precomputed bias (includes causal mask) | |
| out = F.scaled_dot_product_attention(q, k, v, attn_mask=alibi_bias, **_gqa_kw()) | |
| elif cache_k is None or start_pos == 0: | |
| out = F.scaled_dot_product_attention(q, k, v, is_causal=True, **_gqa_kw()) | |
| elif causal_extend and T > 1: | |
| # Verification path: real causal mask within the draft chunk. | |
| # Token at position start_pos+i may attend to cache[0..start_pos+i]. | |
| # [T, S] float mask — tiny for verify-sized T (S≤ few k for short ctx). | |
| S_full = k.shape[2] | |
| qi = torch.arange(T, device=x.device).unsqueeze(1) | |
| kj = torch.arange(S_full, device=x.device).unsqueeze(0) | |
| allow = kj <= (start_pos + qi) # [T, S] | |
| mask = torch.where(allow, 0.0, float("-inf")).to(q.dtype) | |
| out = F.scaled_dot_product_attention( | |
| q, k, v, attn_mask=mask[None, None], **_gqa_kw()) | |
| else: | |
| # inference extend: new tokens attend to all cached + causal among themselves. | |
| # For accept-all (quality not a goal), use is_causal=False to avoid | |
| # materializing a [T, S] mask (which OOMs at S=1M+). SDPA flash kernel | |
| # handles arbitrary S without materializing the attention matrix. | |
| out = F.scaled_dot_product_attention(q, k, v, is_causal=False, **_gqa_kw()) | |
| out = out.transpose(1, 2).reshape(B, T, -1) # [B, T, n_heads*hd] | |
| return self.wo(out) | |
| # ------------------------------------------------------------------ | |
| # MoBA: Mixture of Block Attention (block-sparse) | |
| # ------------------------------------------------------------------ | |
| def forward_moba(self, x, cache_k, cache_v, k_bar, n_filled_blocks, | |
| current_block_fill): | |
| """Block-sparse attention via MoBA (vectorized, accept-all optimized). | |
| x: [B, T, d] — new tokens (T = K+1 during generation) | |
| cache_k/v: [B, n_blocks, block_size, n_kv, hd] — block-structured KV cache | |
| k_bar: [B, n_blocks, n_kv, hd] — precomputed mean-pooled K per block | |
| n_filled_blocks: int — how many complete blocks are in the cache | |
| current_block_fill: int — how many tokens in the current (partial) block | |
| For accept-all: all queries attend to the same top-k blocks | |
| (most popular across all query tokens + heads). This avoids per-query | |
| routing complexity and is CUDA-graph compatible with fixed block indices. | |
| Returns: [B, T, d] | |
| """ | |
| B, T, _ = x.shape | |
| BS = self.block_size | |
| n_avail = n_filled_blocks + 1 # filled blocks + current partial | |
| k_top = min(self.top_k, n_avail) | |
| q = self.wq(x).view(B, T, self.n_heads, self.hd) | |
| k_new = self.wk(x).view(B, T, self.n_kv, self.hd) | |
| v_new = self.wv(x).view(B, T, self.n_kv, self.hd) | |
| # Write new K,V into the current block | |
| blk_idx = n_filled_blocks | |
| cache_k[:, blk_idx, current_block_fill:current_block_fill + T] = k_new | |
| cache_v[:, blk_idx, current_block_fill:current_block_fill + T] = v_new | |
| # Update k_bar for the current block (mean of filled portion) | |
| new_count = current_block_fill + T | |
| if new_count <= BS: | |
| filled_k = cache_k[:, blk_idx, :new_count] # [B, new_count, n_kv, hd] | |
| k_bar[:, blk_idx] = filled_k.mean(dim=1) | |
| # --- Gating (vectorized) --- | |
| q_t = q.transpose(1, 2) # [B, n_heads, T, hd] | |
| kb = k_bar[:, :n_avail] # [B, n_avail, n_kv, hd] | |
| # GQA expand: [B, n_avail, n_heads, hd] | |
| rep = self.n_heads // self.n_kv | |
| kb_exp = kb.unsqueeze(2).expand(-1, -1, rep, -1, -1).reshape( | |
| B, n_avail, self.n_heads, self.hd) | |
| kb_exp = kb_exp.transpose(1, 2) # [B, n_heads, n_avail, hd] | |
| # Gating scores: [B, n_heads, T, n_avail] | |
| scores = torch.einsum('bhtd,bhnd->bhtn', q_t, kb_exp) | |
| # Causal: mask out future blocks | |
| if n_avail < scores.shape[-1]: | |
| scores[:, :, :, n_avail:] = float('-inf') | |
| # Accept-all: average scores over heads + tokens → [B, n_avail] | |
| avg_scores = scores.mean(dim=(1, 2)) | |
| # Always include current block; select top-k | |
| popular = avg_scores.topk(k_top, dim=-1).indices # [B, k_top] | |
| # --- Block attention (k_top SDPA calls) --- | |
| outputs = [] | |
| for ki in range(k_top): | |
| blk = popular[:, ki] # [B] — block index | |
| # Gather K,V for this block: [B, BS, n_kv, hd] | |
| if B == 1: | |
| bi = int(blk.item()) | |
| bk = cache_k[:, bi, :BS] | |
| bv = cache_v[:, bi, :BS] | |
| else: | |
| bk = cache_k[torch.arange(B), blk, :BS] | |
| bv = cache_v[torch.arange(B), blk, :BS] | |
| bk = bk.transpose(1, 2) # [B, n_kv, BS, hd] | |
| bv = bv.transpose(1, 2) | |
| bk, bv = _expand_kv(q_t, bk, bv) | |
| # Current block: causal; historical: non-causal | |
| if ki == k_top - 1: # current block is usually highest score | |
| out_i = F.scaled_dot_product_attention( | |
| q_t, bk, bv, is_causal=True, **_gqa_kw()) | |
| else: | |
| out_i = F.scaled_dot_product_attention( | |
| q_t, bk, bv, is_causal=False, **_gqa_kw()) | |
| outputs.append(out_i) | |
| # Combine (equal weights — accept-all, quality not a goal) | |
| out = torch.stack(outputs, dim=0).mean(dim=0) # [B, n_heads, T, hd] | |
| out = out.transpose(1, 2).reshape(B, T, -1) | |
| return self.wo(out) | |
| # ------------------------------------------------------------------ | |
| # v5: Compressed MoBA — KV compression + FP8 cache + within-block sparse | |
| # ------------------------------------------------------------------ | |
| def compress_tokens(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: | |
| """DeepSeek V4-style KV compression. | |
| x: [B, T, d] — raw hidden states for T tokens | |
| Returns: (comp_kv, comp_weights) where | |
| comp_kv: [B, T, n_kv*hd] — uncompressed KV entries (one per token) | |
| comp_weights:[B, T, n_kv*hd] — softmax weights for pooling | |
| Compression happens in forward_moba_v5 by pooling every m entries. | |
| """ | |
| comp_k = self.w_kv_comp(x) # [B, T, n_kv*hd] | |
| comp_v = self.w_v_comp(x) # [B, T, n_kv*hd] | |
| comp_z = self.w_z_comp(x) # [B, T, n_kv*hd] | |
| return comp_k, comp_v, comp_z | |
| def pool_compressed(self, comp_k: torch.Tensor, comp_v: torch.Tensor, | |
| comp_z: torch.Tensor, m: int): | |
| """Pool every m compressed entries into 1 via learned weighted softmax. | |
| Returns: (pooled_k, pooled_v) each [B, T//m, n_kv*hd] — the pooling | |
| weights come from comp_z (shared for K and V so an entry's K and V | |
| describe the same token group). | |
| """ | |
| B, T, D = comp_k.shape | |
| n_blocks = T // m | |
| comp_k = comp_k[:, :n_blocks * m].view(B, n_blocks, m, D) | |
| comp_v = comp_v[:, :n_blocks * m].view(B, n_blocks, m, D) | |
| comp_z = comp_z[:, :n_blocks * m].view(B, n_blocks, m, D) | |
| # Add positional bias [m, D] broadcast to [1, 1, m, D] | |
| z = comp_z + self.comp_bias.unsqueeze(0).unsqueeze(0) | |
| weights = F.softmax(z, dim=2) # [B, n_blocks, m, D] | |
| pooled_k = (weights * comp_k).sum(dim=2) | |
| pooled_v = (weights * comp_v).sum(dim=2) | |
| return pooled_k, pooled_v | |
| def forward_moba_v5(self, x: torch.Tensor, | |
| cache_k, cache_v, k_bar_comp, | |
| n_filled_blocks: int, current_block_fill: int, | |
| ring_k=None, ring_v=None, ring_len: int = 0, | |
| cos: torch.Tensor = None, sin: torch.Tensor = None, | |
| base_pos: int = 0): | |
| """v5 block-sparse attention: compressed far memory + raw recent window. | |
| KV sources, in order: | |
| - top-k FILLED compressed blocks (4x pooled, FP8, strided) — old ctx | |
| - raw-KV ring: last cfg.raw_window forwarded tokens, exact K,V | |
| - this chunk's raw K,V under a real causal mask | |
| cache_k/cache_v: [B, n_blocks, BS_comp, n_kv, hd] (FP8 or bf16) | |
| ring_k/ring_v: [B, W, n_kv, hd] bf16 — ring[:ring_len] is valid; | |
| this step's K,V are appended in place AFTER attention. | |
| Returns: [B, T, d] | |
| """ | |
| B, T, _ = x.shape | |
| m = self.compress_m | |
| BS_comp = self.cfg.block_size_comp | |
| k_top = self.top_k | |
| stride = self.sparse_stride | |
| W = ring_k.shape[1] if ring_k is not None else 0 | |
| # --- Step 1-2: compress + pool new tokens (separate K and V) --- | |
| comp_k, comp_v, comp_z = self.compress_tokens(x) | |
| pooled_k, pooled_v = self.pool_compressed(comp_k, comp_v, comp_z, m) | |
| n_new = pooled_k.shape[1] | |
| pooled_k = pooled_k.view(B, n_new, self.n_kv, self.hd) | |
| pooled_v = pooled_v.view(B, n_new, self.n_kv, self.hd) | |
| # RoPE on pooled keys at group-center position (rotate after pooling | |
| # to avoid averaging differently-rotated keys). | |
| if cos is not None and cos.numel() > 0: | |
| ent_idx = base_pos // m + torch.arange(n_new, device=x.device) | |
| key_pos = (ent_idx * m + m // 2).clamp(max=cos.shape[0] - 1) | |
| pooled_k = apply_rope(pooled_k, cos[key_pos], sin[key_pos]) | |
| # --- Step 3: write compressed entries (block-straddling) --- | |
| pk_store = pooled_k.to(torch.float8_e4m3fn) if self.use_fp8 else pooled_k | |
| pv_store = pooled_v.to(torch.float8_e4m3fn) if self.use_fp8 else pooled_v | |
| remaining = n_new | |
| write_offset = 0 | |
| cur_fill = current_block_fill | |
| cur_blk = n_filled_blocks | |
| while remaining > 0: | |
| space = BS_comp - cur_fill | |
| if space <= 0: | |
| cur_blk += 1 | |
| cur_fill = 0 | |
| space = BS_comp | |
| to_write = min(remaining, space) | |
| cache_k[:, cur_blk, cur_fill:cur_fill + to_write] = \ | |
| pk_store[:, write_offset:write_offset + to_write] | |
| cache_v[:, cur_blk, cur_fill:cur_fill + to_write] = \ | |
| pv_store[:, write_offset:write_offset + to_write] | |
| write_offset += to_write | |
| remaining -= to_write | |
| cur_fill += to_write | |
| # --- Step 4: update k_bar for touched blocks --- | |
| dt = pooled_k.dtype | |
| n_blocks_touched = cur_blk - n_filled_blocks + 1 | |
| for b in range(n_blocks_touched): | |
| bi = n_filled_blocks + b | |
| if b == 0 and current_block_fill > 0: | |
| count = min(BS_comp, current_block_fill + n_new) if b == n_blocks_touched - 1 else BS_comp | |
| elif b == n_blocks_touched - 1: | |
| count = cur_fill | |
| else: | |
| count = BS_comp | |
| if count > 0 and bi < k_bar_comp.shape[1]: | |
| filled_k = cache_k[:, bi, :count].to(dt) | |
| k_bar_comp[:, bi] = filled_k.mean(dim=1) | |
| # --- Step 5: queries + gating over FILLED blocks (mean-q, exact) --- | |
| q = self.wq(x).view(B, T, self.n_heads, self.hd) | |
| if cos is not None and cos.numel() > 0: | |
| q = apply_rope(q, cos[base_pos:base_pos + T], | |
| sin[base_pos:base_pos + T]) | |
| q_t = q.transpose(1, 2) | |
| n_hist = n_filled_blocks | |
| if n_hist > 0: | |
| kb = k_bar_comp[:, :n_hist] | |
| rep = self.n_heads // self.n_kv | |
| q_mean = q.view(B, T, self.n_kv, rep, self.hd).mean(dim=(1, 3)) | |
| avg_scores = torch.einsum('bkd,bnkd->bkn', q_mean, kb).mean(1) | |
| if n_hist < k_top: | |
| padding = torch.full((B, k_top - n_hist), float('-inf'), | |
| device=avg_scores.device, | |
| dtype=avg_scores.dtype) | |
| avg_scores = torch.cat([avg_scores, padding], dim=-1) | |
| popular = avg_scores.topk(k_top, dim=-1).indices | |
| sel_valid = popular < n_hist | |
| # --- Step 6: gather KV = [compressed past ; ring ; new chunk] --- | |
| parts_k, parts_v, parts_valid = [], [], [] | |
| if n_hist > 0: | |
| ck = cache_k.view(torch.uint8) if self.use_fp8 else cache_k | |
| cv = cache_v.view(torch.uint8) if self.use_fp8 else cache_v | |
| if B == 1: | |
| bi = popular[0] | |
| gk = ck[:, bi.clamp(max=n_hist - 1), :BS_comp] | |
| gv = cv[:, bi.clamp(max=n_hist - 1), :BS_comp] | |
| col_valid = sel_valid[0].repeat_interleave(BS_comp) | |
| else: | |
| bi = popular.clamp(max=n_hist - 1) | |
| ar = torch.arange(B, device=x.device).unsqueeze(1) | |
| gk = ck[ar, bi, :BS_comp] | |
| gv = cv[ar, bi, :BS_comp] | |
| col_valid = sel_valid.repeat_interleave(BS_comp, dim=1) | |
| gk = gk.reshape(B, k_top * BS_comp, self.n_kv, self.hd) | |
| gv = gv.reshape(B, k_top * BS_comp, self.n_kv, self.hd) | |
| if self.use_fp8: | |
| gk = gk.view(torch.float8_e4m3fn).to(q_t.dtype) | |
| gv = gv.view(torch.float8_e4m3fn).to(q_t.dtype) | |
| if stride > 1: | |
| gk, gv = gk[:, ::stride], gv[:, ::stride] | |
| col_valid = col_valid[::stride] if col_valid.dim() == 1 \ | |
| else col_valid[:, ::stride] | |
| parts_k.append(gk) | |
| parts_v.append(gv) | |
| parts_valid.append(col_valid.unsqueeze(0) if col_valid.dim() == 1 | |
| else col_valid) | |
| if ring_k is not None and ring_len > 0: | |
| parts_k.append(ring_k[:, :ring_len]) | |
| parts_v.append(ring_v[:, :ring_len]) | |
| parts_valid.append(torch.ones(B, ring_len, dtype=torch.bool, | |
| device=x.device)) | |
| k_new = self.wk(x).view(B, T, self.n_kv, self.hd) | |
| v_new = self.wv(x).view(B, T, self.n_kv, self.hd) | |
| if cos is not None and cos.numel() > 0: | |
| k_new = apply_rope(k_new, cos[base_pos:base_pos + T], | |
| sin[base_pos:base_pos + T]) | |
| # --- Step 7: single SDPA over [past ; new] with combined mask --- | |
| kk = torch.cat(parts_k + [k_new], dim=1).transpose(1, 2) | |
| vv = torch.cat(parts_v + [v_new], dim=1).transpose(1, 2) | |
| kk, vv = _expand_kv(q_t, kk, vv) | |
| if parts_valid: | |
| S_past = kk.shape[2] - T | |
| causal = torch.ones(T, T, dtype=torch.bool, | |
| device=x.device).tril_() | |
| pv = torch.cat(parts_valid, dim=1) | |
| mask = torch.cat([ | |
| pv.unsqueeze(1).expand(B, T, S_past), | |
| causal.unsqueeze(0).expand(B, T, T)], dim=2).unsqueeze(1) | |
| out = F.scaled_dot_product_attention(q_t, kk, vv, attn_mask=mask, | |
| **_gqa_kw()) | |
| else: | |
| out = F.scaled_dot_product_attention(q_t, kk, vv, is_causal=True, | |
| **_gqa_kw()) | |
| # --- Step 8: append this chunk's raw K,V to the ring --- | |
| if ring_k is not None: | |
| if T >= W: | |
| ring_k[:] = k_new[:, T - W:] | |
| ring_v[:] = v_new[:, T - W:] | |
| elif ring_len + T <= W: | |
| ring_k[:, ring_len:ring_len + T] = k_new | |
| ring_v[:, ring_len:ring_len + T] = v_new | |
| else: | |
| keep = W - T | |
| ring_k[:, :keep] = ring_k[:, ring_len - keep:ring_len].clone() | |
| ring_v[:, :keep] = ring_v[:, ring_len - keep:ring_len].clone() | |
| ring_k[:, keep:W] = k_new | |
| ring_v[:, keep:W] = v_new | |
| out = out.transpose(1, 2).reshape(B, T, -1) | |
| return self.wo(out) | |
| # -------------------------------------------------------------------------------------- | |
| # SwiGLU feed-forward | |
| # -------------------------------------------------------------------------------------- | |
| class FeedForward(nn.Module): | |
| def __init__(self, cfg: Config): | |
| super().__init__() | |
| d, inter = cfg.d_model, cfg.intermediate | |
| self.w_gate = nn.Linear(d, inter, bias=False) | |
| self.w_up = nn.Linear(d, inter, bias=False) | |
| self.w_down = nn.Linear(inter, d, bias=False) | |
| def forward(self, x): | |
| return self.w_down(F.silu(self.w_gate(x)) * self.w_up(x)) | |
| # -------------------------------------------------------------------------------------- | |
| # Sequential draft block (EAGLE-lite): a single small transformer layer that | |
| # evolves its own hidden state across the draft chain. Unlike the flat heads, | |
| # each draft step sees the full drafted prefix via causal attention. | |
| # -------------------------------------------------------------------------------------- | |
| class SeqDraftBlock(nn.Module): | |
| def __init__(self, cfg: Config): | |
| super().__init__() | |
| d = cfg.d_model | |
| nh, hd = cfg.n_heads, cfg.head_dim | |
| self.nh, self.hd = nh, hd | |
| inter = int(d * cfg.seq_ffn_mult) | |
| self.norm1 = nn.RMSNorm(d) | |
| self.wq = nn.Linear(d, nh * hd, bias=False) | |
| self.wk = nn.Linear(d, nh * hd, bias=False) | |
| self.wv = nn.Linear(d, nh * hd, bias=False) | |
| self.wo = nn.Linear(nh * hd, d, bias=False) | |
| self.norm2 = nn.RMSNorm(d) | |
| self.w_gate = nn.Linear(d, inter, bias=False) | |
| self.w_up = nn.Linear(d, inter, bias=False) | |
| self.w_down = nn.Linear(inter, d, bias=False) | |
| def forward(self, x): | |
| # x [B,T,d] — plain causal self-attention over the draft prefix | |
| B, T, _ = x.shape | |
| h = self.norm1(x) | |
| q = self.wq(h).view(B, T, self.nh, self.hd).transpose(1, 2) | |
| k = self.wk(h).view(B, T, self.nh, self.hd).transpose(1, 2) | |
| v = self.wv(h).view(B, T, self.nh, self.hd).transpose(1, 2) | |
| o = F.scaled_dot_product_attention(q, k, v, is_causal=True) | |
| x = x + self.wo(o.transpose(1, 2).reshape(B, T, -1)) | |
| h2 = self.norm2(x) | |
| x = x + self.w_down(F.silu(self.w_gate(h2)) * self.w_up(h2)) | |
| return x | |
| # -------------------------------------------------------------------------------------- | |
| # Block | |
| # -------------------------------------------------------------------------------------- | |
| class Block(nn.Module): | |
| def __init__(self, cfg: Config): | |
| super().__init__() | |
| self.norm1 = nn.RMSNorm(cfg.d_model) | |
| self.attn = Attention(cfg) | |
| self.norm2 = nn.RMSNorm(cfg.d_model) | |
| self.ffn = FeedForward(cfg) | |
| def forward(self, x, cos, sin, cache_k=None, cache_v=None, start_pos=0, | |
| alibi_bias=None, sliding_window=False, causal_extend=False): | |
| x = x + self.attn(self.norm1(x), cos, sin, cache_k, cache_v, start_pos, | |
| alibi_bias, sliding_window, causal_extend) | |
| x = x + self.ffn(self.norm2(x)) | |
| return x | |
| def forward_moba(self, x, cache_k, cache_v, k_bar, n_filled_blocks, | |
| current_block_fill): | |
| x = x + self.attn.forward_moba(self.norm1(x), cache_k, cache_v, k_bar, | |
| n_filled_blocks, current_block_fill) | |
| x = x + self.ffn(self.norm2(x)) | |
| return x | |
| def forward_moba_v5(self, x, cache_k, cache_v, k_bar_comp, n_filled_blocks, | |
| current_block_fill, ring_k=None, ring_v=None, | |
| ring_len=0, cos=None, sin=None, base_pos=0): | |
| x = x + self.attn.forward_moba_v5(self.norm1(x), cache_k, cache_v, | |
| k_bar_comp, n_filled_blocks, | |
| current_block_fill, | |
| ring_k, ring_v, ring_len, | |
| cos, sin, base_pos) | |
| if hasattr(self, '_compiled_ffn_fn'): | |
| x = x + self._compiled_ffn_fn(x) | |
| else: | |
| x = x + self.ffn(self.norm2(x)) | |
| return x | |
| # -------------------------------------------------------------------------------------- | |
| # Full model | |
| # -------------------------------------------------------------------------------------- | |
| class SpecModel(nn.Module): | |
| def __init__(self, cfg: Config): | |
| super().__init__() | |
| self.cfg = cfg | |
| self.tok_emb = nn.Embedding(cfg.vocab_size, cfg.d_model) | |
| self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.n_layers)]) | |
| self.norm_f = nn.RMSNorm(cfg.d_model) | |
| # tied output head handled in forward_logits | |
| self.rope_cos: torch.Tensor | |
| self.rope_sin: torch.Tensor | |
| self.register_buffer("rope_cos", torch.empty(0), persistent=False) | |
| self.register_buffer("rope_sin", torch.empty(0), persistent=False) | |
| self._rope_built = False | |
| self._alibi_built = False | |
| self.alibi_bias: torch.Tensor = None | |
| # --- Speculative heads --- | |
| # r == 1 (legacy): per-head low-rank [K,d,1] + [K,1,vocab]; each head can | |
| # only emit argmax/argmin of its fixed weight row (2 tokens total). | |
| # r > 1 (v6 hybrid): conditioned block heads + parallel tail heads, all | |
| # sharing ONE vocab projection so every head emits a real | |
| # context-dependent distribution over the full vocab. | |
| K, d, r, v = cfg.medusa_heads, cfg.d_model, cfg.medusa_rank, cfg.vocab_size | |
| if r == 1: | |
| self.medusa_down_weight = nn.Parameter(torch.empty(K, d, r)) | |
| self.medusa_out_weight = nn.Parameter(torch.empty(K, r, v)) | |
| else: | |
| e, g = cfg.medusa_emb_rank, cfg.medusa_cond_group | |
| kc, kp = cfg.medusa_cond_heads, cfg.medusa_par_heads | |
| self.spec_tok_proj = nn.Parameter(torch.empty(d, e)) | |
| self.spec_cond_w = nn.Parameter(torch.empty(kc, d + g * e, r)) | |
| self.spec_par_w = nn.Parameter(torch.empty(kp, d, r)) | |
| self.spec_vocab = nn.Parameter(torch.empty(r, v)) | |
| self.spec_bias = nn.Parameter(torch.zeros(v)) | |
| # v7 dspark-style: markov head = intra-chain token dependency. | |
| # logits_j += W2(emb(prev predicted token)) — the lightweight | |
| # sequential module that fixes parallel-emission suffix decay. | |
| mr = cfg.medusa_rank | |
| self.spec_markov_emb = nn.Parameter(torch.empty(v, mr)) | |
| self.spec_markov_w2 = nn.Parameter(torch.empty(mr, v)) | |
| # confidence head: P(draft token accepted) -> adaptive verify_k | |
| # input = [h_anchor ; group cond ; slot code] | |
| self.spec_conf_w = nn.Parameter(torch.empty(d + g * e + r, 1)) | |
| # hi-rank near-field heads (capacity fix): heads 0..H-1 get a | |
| # wide code + their own vocab projection instead of sharing the | |
| # r=256 bottleneck. Replaces the low-rank path for those slots. | |
| H = getattr(cfg, "medusa_hi_heads", 0) | |
| if H > 0: | |
| R = cfg.medusa_hi_rank | |
| self.spec_hi_w = nn.Parameter(torch.empty(H, d + g * e, R)) | |
| self.spec_hi_vocab = nn.Parameter(torch.empty(R, v)) | |
| self.spec_hi_bias = nn.Parameter(torch.zeros(v)) | |
| # EAGLE-lite sequential draft module: evolves its own hidden | |
| # state token-by-token — the fix for the flat-head information | |
| # bottleneck. Output via tied tok_emb (zero output params). | |
| if getattr(cfg, "medusa_seq_len", 0) > 0: | |
| # input: [h_i ; emb(tok_i) ; compressed last-G tokens] | |
| self.spec_seq_fc = nn.Linear(2 * d + g * e, d, bias=False) | |
| self.spec_seq_blk = SeqDraftBlock(cfg) | |
| self.reset_parameters() | |
| def reset_parameters(self): | |
| # small init for stability; std ~ 1/sqrt(d) scaled | |
| std = 0.02 | |
| for p in self.parameters(): | |
| if p.dim() > 1: | |
| nn.init.normal_(p, std=std) | |
| # scale embeddings down a touch | |
| nn.init.normal_(self.tok_emb.weight, std=std) | |
| # v7 zero-init: markov bias is a no-op and conf is neutral (0.5) | |
| # until trained — keeps checkpoints/behaviour compatible. | |
| if hasattr(self, "spec_markov_w2"): | |
| nn.init.zeros_(self.spec_markov_w2) | |
| nn.init.zeros_(self.spec_conf_w) | |
| def build_rope(self, device, dtype): | |
| if not self._rope_built: | |
| base = self.cfg.rope_base * self.cfg.rope_ntk_scale | |
| cos, sin = precompute_rope(self.cfg.head_dim, self.cfg.max_seq_len, | |
| base, device, dtype) | |
| self.rope_cos = cos | |
| self.rope_sin = sin | |
| self._rope_built = True | |
| def build_alibi(self, q_len: int, device, dtype): | |
| """Precompute fixed ALiBi bias for sliding window: [n_heads, q_len, W].""" | |
| W = self.cfg.window_size if self.cfg.window_size > 0 else self.cfg.max_seq_len | |
| if not self._alibi_built or self.alibi_bias is None or self.alibi_bias.shape[1] != q_len: | |
| self.alibi_bias = precompute_alibi_bias( | |
| self.cfg.n_heads, q_len, W, dtype, device) | |
| self._alibi_built = True | |
| def forward(self, input_ids, cache_k=None, cache_v=None, start_pos=0, | |
| sliding_window=False, use_checkpoint=False, causal_extend=False): | |
| """ | |
| input_ids: [B, T] | |
| cache_k/v: list per layer of [B, W, n_kv, hd] (None for training) | |
| start_pos: int (cached length, growing cache only) | |
| sliding_window: if True, use fixed-size cache with ALiBi | |
| use_checkpoint: if True, use gradient checkpointing (saves memory, | |
| costs ~30% more compute — allows bigger batches) | |
| causal_extend: if True, causal-mask within the new chunk (speculative | |
| verification — the accept-all path skips this). | |
| Returns hidden states [B, T, d] (pre-output-head). | |
| """ | |
| B, T = input_ids.shape | |
| dt = self.tok_emb.weight.dtype | |
| dev = input_ids.device | |
| alibi = None | |
| if self.cfg.use_alibi and sliding_window: | |
| if self.cfg.window_size == T: | |
| pass | |
| else: | |
| self.build_alibi(T, dev, dt) | |
| alibi = self.alibi_bias | |
| elif not self.cfg.use_alibi: | |
| self.build_rope(dev, dt) | |
| x = self.tok_emb(input_ids) # [B, T, d] | |
| for i, blk in enumerate(self.blocks): | |
| ck = cache_k[i] if cache_k is not None else None | |
| cv = cache_v[i] if cache_v is not None else None | |
| if use_checkpoint and self.training and ck is None: | |
| # Gradient checkpointing: recompute forward in backward | |
| # Saves ~8x activation memory (only store layer outputs) | |
| rope_cos = self.rope_cos | |
| rope_sin = self.rope_sin | |
| def run_blk(x, blk=blk, rope_cos=rope_cos, rope_sin=rope_sin, | |
| alibi=alibi, sliding_window=sliding_window): | |
| return blk(x, rope_cos, rope_sin, None, None, 0, | |
| alibi_bias=alibi, sliding_window=sliding_window) | |
| x = torch.utils.checkpoint.checkpoint(run_blk, x, use_reentrant=False) | |
| else: | |
| x = blk(x, self.rope_cos, self.rope_sin, ck, cv, start_pos, | |
| alibi_bias=alibi, sliding_window=sliding_window, | |
| causal_extend=causal_extend) | |
| return self.norm_f(x) | |
| def forward_moba(self, input_ids, cache_k, cache_v, k_bar_list, | |
| n_filled_blocks, current_block_fill): | |
| """MoBA forward pass: block-sparse attention for long context. | |
| input_ids: [B, T] | |
| cache_k/v: list per layer of [B, n_blocks, block_size, n_kv, hd] | |
| k_bar_list: list per layer of [B, n_blocks, n_kv, hd] (precomputed block means) | |
| n_filled_blocks: int — complete blocks in cache | |
| current_block_fill: int — tokens in current partial block | |
| Returns hidden states [B, T, d] | |
| """ | |
| x = self.tok_emb(input_ids) | |
| for i, blk in enumerate(self.blocks): | |
| x = blk.forward_moba(x, cache_k[i], cache_v[i], k_bar_list[i], | |
| n_filled_blocks, current_block_fill) | |
| return self.norm_f(x) | |
| def forward_moba_v5(self, input_ids, cache_k, cache_v, k_bar_comp, | |
| n_filled_blocks, current_block_fill, | |
| ring_k=None, ring_v=None, ring_len=0): | |
| """v5 MoBA forward: compressed far memory + raw recent window. | |
| input_ids: [B, T] | |
| cache_k/cache_v: per-layer [B, n_blocks, BS_comp, n_kv, hd] (FP8/bf16) | |
| k_bar_comp: per-layer [B, n_blocks, n_kv, hd] — mean pooled K per block | |
| ring_k/ring_v: per-layer [B, raw_window, n_kv, hd] — exact recent KV | |
| ring_len: valid entries in the ring (same across layers) | |
| Returns hidden states [B, T, d] | |
| """ | |
| x = self.tok_emb(input_ids) | |
| cos = sin = None | |
| if not self.cfg.use_alibi: | |
| self.build_rope(input_ids.device, x.dtype) | |
| cos, sin = self.rope_cos, self.rope_sin | |
| base_pos = (n_filled_blocks * self.cfg.block_size_comp | |
| + current_block_fill) * self.cfg.kv_compress_m | |
| for i, blk in enumerate(self.blocks): | |
| x = blk.forward_moba_v5( | |
| x, cache_k[i], cache_v[i], k_bar_comp[i], | |
| n_filled_blocks, current_block_fill, | |
| None if ring_k is None else ring_k[i], | |
| None if ring_v is None else ring_v[i], | |
| ring_len, cos, sin, base_pos) | |
| return self.norm_f(x) | |
| def init_compressor_from_attn(self): | |
| """Initialize the KV compressor so pooled compressed K/V ~= mean-pooled | |
| real K and V. | |
| w_kv_comp <- wk, w_v_comp <- wv | |
| w_z_comp <- 0, comp_bias <- 0 (uniform softmax -> mean pooling) | |
| """ | |
| for blk in self.blocks: | |
| a = blk.attn | |
| if a.compress_m > 1: | |
| a.w_kv_comp.weight.data.copy_(a.wk.weight.data) | |
| a.w_v_comp.weight.data.copy_(a.wv.weight.data) | |
| a.w_z_comp.weight.data.zero_() | |
| a.comp_bias.data.zero_() | |
| def lm_head(self, hidden: torch.Tensor) -> torch.Tensor: | |
| """Tied output projection: hidden [., d] -> [., vocab].""" | |
| return F.linear(hidden, self.tok_emb.weight) | |
| def medusa_logits(self, hidden_last: torch.Tensor): | |
| """hidden_last: [B, d] (single position). Returns [B, K, vocab] logits. | |
| Legacy r=1 path only — kept for old checkpoints.""" | |
| down = torch.einsum('bd,kdr->bkr', hidden_last, self.medusa_down_weight) | |
| return torch.einsum('bkr,krv->bkv', down, self.medusa_out_weight) | |
| def medusa_argmax(self, hidden_last: torch.Tensor): | |
| """Greedy tokens from each Medusa head. Returns [B, K] long tensor. | |
| Legacy r=1 path.""" | |
| logits = self.medusa_logits(hidden_last) # [B, K, vocab] | |
| return logits.argmax(-1) # [B, K] | |
| # ------------------------------------------------------------------ | |
| # Precomputed argmax mode (r=1 only): skip the 410MB medusa out einsum | |
| # ------------------------------------------------------------------ | |
| def precompute_medusa_tokens(self): | |
| """For r=1: argmax(s * w) = argmax(w) if s>0, argmin(w) if s<0. | |
| Precompute argmax and argmin of each head's out_weight row once. | |
| After this, medusa_argmax_fast needs only the down projection + sign check.""" | |
| assert self.cfg.medusa_rank == 1, "precomputed mode requires r=1" | |
| w = self.medusa_out_weight.squeeze(1) # [K, vocab] | |
| self._medusa_argmax = w.argmax(dim=-1) # [K] | |
| self._medusa_argmin = w.argmin(dim=-1) # [K] | |
| self._medusa_precomputed = True | |
| def medusa_argmax_fast(self, hidden_last: torch.Tensor): | |
| """r=1 fast path: down projection + sign check, no medusa out einsum. | |
| Returns [B, K] long tensor.""" | |
| # down: [B, K] — just a matmul [B, d] x [d, K] (squeeze the r=1 dim) | |
| down = torch.einsum('bd,kd->bk', hidden_last, | |
| self.medusa_down_weight.squeeze(-1)) # [B, K] | |
| positive = (down > 0).long() # [B, K] | |
| # where positive: argmax, else argmin | |
| return torch.where(positive.bool(), | |
| self._medusa_argmax.expand_as(down), | |
| self._medusa_argmin.expand_as(down)) # [B, K] | |
| # ------------------------------------------------------------------ | |
| # v6 hybrid heads: conditioned blocks + parallel tail, shared vocab | |
| # ------------------------------------------------------------------ | |
| def spec_cond_logits(self, h_last: torch.Tensor, cond_embs: torch.Tensor, | |
| head_idx: torch.Tensor): | |
| """Conditioned-head logits for chosen heads. | |
| h_last: [B, d] — anchor hidden state | |
| cond_embs: [B, G, e] — compressed embs of the previous group's tokens | |
| head_idx: [S] — which conditioned heads to evaluate | |
| Returns [B, S, V] logits.""" | |
| B = h_last.shape[0] | |
| cond_flat = cond_embs.reshape(B, -1) # [B, G*e] | |
| inp = torch.cat([h_last, cond_flat], dim=-1) # [B, d+G*e] | |
| codes = F.silu(torch.einsum('be,ser->bsr', inp, | |
| self.spec_cond_w[head_idx])) # [B,S,r] | |
| return codes @ self.spec_vocab + self.spec_bias # [B,S,V] | |
| def spec_par_logits(self, h_last: torch.Tensor, | |
| head_idx: torch.Tensor = None): | |
| """Parallel-tail logits. h_last [B,d] -> [B,Kp,V] (or [B,S,V] if subsampled).""" | |
| w = self.spec_par_w if head_idx is None else self.spec_par_w[head_idx] | |
| codes = F.silu(torch.einsum('bd,sdr->bsr', h_last, w)) # [B,S,r] | |
| return codes @ self.spec_vocab + self.spec_bias # [B,S,V] | |
| def spec_draft(self, h_anchor: torch.Tensor, hist_ids: torch.Tensor, | |
| first_head: int = 0, pending_id: int = None, | |
| prefix_ids: torch.Tensor = None, | |
| n_cond: int = None, n_par: int = None, | |
| return_conf: bool = False, | |
| sample: bool = False, temperature: float = 1.0, | |
| top_p: float = 0.9): | |
| """Draft a token chain from an anchor hidden state. | |
| h_anchor: [B, d] hidden at the last *forwarded* position. | |
| hist_ids: [B, G] last G committed tokens (left-pad with any id). | |
| first_head: number of leading head slots to drop — the committed | |
| tokens already occupying those positions (pending/queue). | |
| pending_id: the committed-but-unforwarded token (first_head=1). | |
| prefix_ids: [B, P<=G] committed tokens forced into the first P chain | |
| slots (generalizes pending_id for a queue of P tokens). | |
| Returns [B, K - first_head] draft token ids. | |
| """ | |
| cfg = self.cfg | |
| B = h_anchor.shape[0] | |
| G, Kc = cfg.medusa_cond_group, cfg.medusa_cond_heads | |
| n_groups = Kc // G | |
| n_cond = n_cond or Kc | |
| n_par = n_par if n_par is not None else cfg.medusa_par_heads | |
| dev = h_anchor.device | |
| if pending_id is not None and prefix_ids is None: | |
| prefix_ids = torch.full((B, 1), pending_id, dtype=torch.long, | |
| device=dev) | |
| n_prefix = 0 if prefix_ids is None else prefix_ids.shape[1] | |
| def _pick(lg): # [B,V] -> [B] | |
| """argmax, or temperature+top-p sample when sample=True.""" | |
| if not sample or temperature <= 0: | |
| return lg.argmax(-1) | |
| lg = lg.float() / temperature | |
| if top_p < 1.0: | |
| s, si = lg.sort(-1, descending=True) | |
| cum = s.softmax(-1).cumsum(-1) | |
| rm = cum - s.softmax(-1) >= top_p # keep cum<p | |
| s = s.masked_fill(rm, float("-inf")) | |
| return si.gather(-1, torch.multinomial( | |
| s.softmax(-1), 1)).squeeze(-1) | |
| return torch.multinomial(lg.softmax(-1), 1).squeeze(-1) | |
| # EAGLE-lite sequential draft path: replaces the conditioned chain | |
| # when the seq module exists — the hidden state evolves per token. | |
| if getattr(cfg, "medusa_seq_len", 0) > 0 and \ | |
| hasattr(self, "spec_seq_fc"): | |
| seq_n = max(0, n_cond - n_prefix) | |
| seq_out = self.spec_draft_seq( | |
| h_anchor, prefix_ids=prefix_ids, hist_ids=hist_ids, | |
| n=seq_n, sample=sample, temperature=temperature, | |
| top_p=top_p) | |
| if n_prefix == 0: | |
| # slot +1 = base's own next-token at the anchor | |
| first = (h_anchor @ self.tok_emb.weight.T | |
| ).argmax(-1, keepdim=True) | |
| seq_out = torch.cat([first, seq_out[:, :-1]], dim=1) | |
| conf_out = torch.full((B, seq_out.shape[1]), 0.5, device=dev) | |
| par_logits = self.spec_par_logits(h_anchor) | |
| if sample: | |
| par_out = torch.stack( | |
| [_pick(par_logits[:, j]) for j in range(n_par)], 1) | |
| else: | |
| par_out = par_logits.argmax(-1)[:, :n_par] | |
| out = torch.cat([seq_out, par_out], dim=1) | |
| if return_conf: | |
| return out, conf_out | |
| return out | |
| use_markov = getattr(self.cfg, "use_markov_head", False) | |
| cond = self.tok_emb(hist_ids) @ self.spec_tok_proj # [B, G, e] | |
| prev = hist_ids[:, -1] # anchor token | |
| preds = [] | |
| confs = [] | |
| for g in range(n_cond // G): | |
| idx = torch.arange(g * G, g * G + G, device=dev) | |
| cond_flat = cond.reshape(B, -1) # [B, G*e] | |
| inp = torch.cat([h_anchor, cond_flat], dim=-1) # [B, d+G*e] | |
| codes = F.silu(torch.einsum('be,ser->bsr', inp, | |
| self.spec_cond_w[idx])) # [B,G,r] | |
| base = codes @ self.spec_vocab + self.spec_bias # [B,G,V] | |
| H = getattr(cfg, "medusa_hi_heads", 0) | |
| if g == 0 and H > 0: | |
| # hi-rank slots: wide code + dedicated vocab proj | |
| codes_hi = F.silu(torch.einsum( | |
| 'be,ser->bsr', inp, self.spec_hi_w[:H])) # [B,H,R] | |
| base = torch.cat( | |
| [codes_hi @ self.spec_hi_vocab + self.spec_hi_bias, | |
| base[:, H:]], dim=1) | |
| codes = torch.cat([codes_hi[..., :codes.shape[-1]], | |
| codes[:, H:]], dim=1) | |
| # per-slot confidence: sigmoid([inp ; code_j] @ conf_w) | |
| cin = torch.cat([inp.unsqueeze(1).expand(B, G, -1), | |
| codes], dim=-1) # [B,G,d+Ge+r] | |
| confs.append(torch.sigmoid(cin @ self.spec_conf_w).squeeze(-1)) | |
| if use_markov: | |
| # DSpark-style intra-block dependency: sequential unroll, | |
| # logits_j += W2(emb(tok_{j-1})) | |
| toks = [] | |
| for j in range(G): | |
| lg = base[:, j] + (self.spec_markov_emb[prev] | |
| # [B,V] | |
| if g == 0 and j < n_prefix: | |
| t = prefix_ids[:, j] | |
| else: | |
| t = _pick(lg) | |
| toks.append(t) | |
| prev = t | |
| tok = torch.stack(toks, dim=1) # [B,G] | |
| else: | |
| tok = (torch.stack([_pick(base[:, j]) for j in range(G)], 1) | |
| if sample else base.argmax(-1)) # [B,G] | |
| if g == 0 and n_prefix: | |
| tok = tok.clone() | |
| tok[:, :n_prefix] = prefix_ids # committed slots | |
| prev = tok[:, -1] | |
| preds.append(tok) | |
| cond = self.tok_emb(tok) @ self.spec_tok_proj # next group cond | |
| cond_out = torch.stack(preds, dim=1).reshape(B, -1) # [B, Kc] | |
| conf_out = torch.stack(confs, dim=1).reshape(B, -1) # [B, Kc] | |
| cond_out = cond_out[:, first_head:] # drop prefix slots | |
| conf_out = conf_out[:, first_head:] | |
| par_logits = self.spec_par_logits(h_anchor) # [B, Kp, V] | |
| if sample: | |
| par_out = torch.stack( | |
| [_pick(par_logits[:, j]) for j in range(n_par)], 1) | |
| else: | |
| par_out = par_logits.argmax(-1)[:, :n_par] # [B, Kp] | |
| out = torch.cat([cond_out, par_out], dim=1) # [B, K-first_head] | |
| if return_conf: | |
| return out, conf_out | |
| return out | |
| def spec_draft_seq(self, h_anchor: torch.Tensor, | |
| prefix_ids: torch.Tensor = None, | |
| hist_ids: torch.Tensor = None, | |
| n: int = 32, | |
| sample: bool = False, temperature: float = 1.0, | |
| top_p: float = 0.9): | |
| """EAGLE-lite sequential draft: the draft module's own hidden state | |
| evolves token-by-token — each step sees the drafted prefix via | |
| causal attention AND the sliding last-G-token conditioning window. | |
| h_anchor: [B, d] hidden at the last *forwarded* position. | |
| prefix_ids: [B, P] committed-but-unforwarded tokens — fed into the | |
| chain first (the seq module processes them as inputs before | |
| emitting predictions). | |
| hist_ids: [B, G] last G committed tokens (pre-prefix) — the sliding | |
| conditioning window. Falls back to the prefix/pad if absent. | |
| Returns [B, n] draft token ids for positions t+2 .. t+1+n | |
| (t+1 is covered by the pending/prefix slot). | |
| """ | |
| B = h_anchor.shape[0] | |
| dev = h_anchor.device | |
| G = self.cfg.medusa_cond_group | |
| n_prefix = 0 if prefix_ids is None else prefix_ids.shape[1] | |
| def _pick(lg): | |
| if not sample or temperature <= 0: | |
| return lg.argmax(-1) | |
| lg = lg.float() / temperature | |
| if top_p < 1.0: | |
| s, si = lg.sort(-1, descending=True) | |
| cum = s.softmax(-1).cumsum(-1) | |
| rm = cum - s.softmax(-1) >= top_p | |
| s = s.masked_fill(rm, float("-inf")) | |
| return si.gather(-1, torch.multinomial( | |
| s.softmax(-1), 1)).squeeze(-1) | |
| return torch.multinomial(lg.softmax(-1), 1).squeeze(-1) | |
| E = self.tok_emb.weight # tied [V,d] | |
| # sliding conditioning window: last G tokens (committed + predicted) | |
| if hist_ids is None: | |
| hist_ids = (prefix_ids[:, :G] if n_prefix >= G | |
| else torch.zeros(B, G, dtype=torch.long, | |
| device=dev)) | |
| window = hist_ids[:, -G:].clone() # [B,G] | |
| h_prev = h_anchor | |
| xs = [] | |
| preds = [] | |
| n_steps = n_prefix + n | |
| for i in range(n_steps): | |
| if i < n_prefix: | |
| tok_in = prefix_ids[:, i] | |
| elif i == 0: | |
| tok_in = (h_anchor @ E.T).argmax(-1) # base's own first token | |
| else: | |
| tok_in = preds[-1] | |
| window = torch.cat([window[:, 1:], tok_in.unsqueeze(1)], 1) | |
| cond = (self.tok_emb(window) @ self.spec_tok_proj | |
| ).reshape(B, -1) # [B,G*e] | |
| x = self.spec_seq_fc(torch.cat( | |
| [h_prev, self.tok_emb(tok_in), cond], dim=-1)) # [B,d] | |
| xs.append(x.unsqueeze(1)) | |
| h_all = self.spec_seq_blk(torch.cat(xs, dim=1)) # [B,i+1,d] | |
| h_prev = h_all[:, -1] # [B,d] | |
| if i >= n_prefix - 1: | |
| lg = h_prev @ E.T # [B,V] | |
| preds.append(_pick(lg)) | |
| return torch.stack(preds, dim=1)[:, :n] | |
| def build_model(cfg: Config, device="cuda") -> SpecModel: | |
| m = SpecModel(cfg).to(device) | |
| dt = getattr(torch, cfg.dtype) | |
| m = m.to(dt) | |
| return m | |