"""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) @torch.no_grad() 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) @torch.no_grad() 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 # ------------------------------------------------------------------ @torch.no_grad() 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 @torch.no_grad() 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] @torch.no_grad() 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