spec100m / code /model.py
Akahsizrr's picture
squash to reclaim LFS quota
a8f07a3
Raw History Blame Contribute Delete
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)
@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<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]
@ self.spec_markov_w2) # [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
@torch.no_grad()
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