spec100m / code /config.py
Akahsizrr's picture
squash to reclaim LFS quota
a8f07a3
Raw History Blame Contribute Delete
10.8 kB
"""Model configuration + parameter counter for the ~100M speculative-decoding model."""
from dataclasses import dataclass, field
@dataclass
class Config:
# --- Tokenizer / vocab ---
vocab_size: int = 50257 # tiktoken GPT-2 BPE; smaller vocab = faster output head
# --- Transformer core (target ~100M params) ---
d_model: int = 384
n_layers: int = 1 # 1 layer: speed ceiling ~65-125k tok/s (memory-bound, not compute)
n_heads: int = 12 # query heads (head_dim = d/n_heads = 32)
n_kv_heads: int = 4 # GQA: 3:1 grouped query attention -> smaller KV cache + faster attn
ffn_mult: float = 2.667 # intermediate = round(ffn_mult * d_model) -> ~2048 (SwiGLU)
rope_base: float = 10000.0
rope_ntk_scale: float = 1.0 # >1: NTK-style base scaling for long-ctx
# extrapolation (base_eff = base * scale)
tie_embeddings: bool = True # share input embedding <-> output projection (saves params + bandwidth)
max_seq_len: int = 65536
# --- Position encoding ---
use_alibi: bool = True # ALiBi (fixed bias, no rotation) for fast inference; False = RoPE
# Sliding window: fixed-size KV cache for constant-time attention + CUDA graph capture.
# 0 = growing cache (old behavior); >0 = fixed window of that size.
# Set to K+1 (= medusa_heads+1) for "no memory" mode: is_causal=True fast path, no mask needed.
window_size: int = 4097
# --- Speculative heads (v6 hybrid: conditioned blocks + parallel tail) ---
# Two head families, all sharing ONE vocab projection (the big win vs
# per-head [r,V] which let each head emit only 2 fixed tokens at r=1):
# - conditioned heads: Kc heads in groups of G. Head k predicts offset k+1
# from [h_t ; compressed embs of the PREVIOUS group's G tokens].
# Sequential over Kc/G groups at inference, teacher-forced in training.
# - parallel heads: Kp heads predict offsets Kc+1..Kc+Kp from h_t alone.
# Shared pieces: spec_vocab [r,V], spec_bias [V], spec_tok_proj [d,e].
medusa_heads: int = 512 # total speculative heads (Kc + Kp)
medusa_cond_heads: int = 256 # conditioned heads (offsets 1..Kc)
medusa_cond_group: int = 8 # heads per conditioning block
medusa_rank: int = 256 # head bottleneck r (1 = legacy argmax-trick mode)
medusa_emb_rank: int = 32 # e: compressed token emb for conditioning
medusa_hidden: int = 0 # 0 = linear low-rank; >0 adds a hidden layer
medusa_hi_heads: int = 0 # first H conditioned heads use hi-rank path
# (own wide vocab proj) — capacity fix for
# near-field acceptance. H <= G.
medusa_hi_rank: int = 1024 # hi-path code dimension
medusa_seq_len: int = 0 # >0: EAGLE-lite sequential draft module —
# a mini transformer block that evolves its
# own hidden state token-by-token (the fix
# for the flat-head information bottleneck)
seq_ffn_mult: float = 4.0 # draft block FFN width (d * mult)
use_markov_head: bool = False # dspark-style intra-chain token dependency;
# sequential unroll — enable once trained
# --- Inference / speed ---
dtype: str = "bfloat16" # Ampere supports bf16 natively -> halves memory bandwidth
compile_mode: str = "default" # torch.compile mode; "reduce-overhead" = CUDA graphs (needs static shapes)
# --- MoBA (Mixture of Block Attention) ---
# Block-sparse attention for long context. KV cache is split into blocks;
# each query attends to top-k blocks selected by Q @ mean(K_block).
# 0 = disabled (use sliding window or growing cache).
moba_block_size: int = 4096 # tokens per block
moba_top_k: int = 3 # blocks each query attends to
# --- v5: KV compression + FP8 cache + within-block sparse ---
# DeepSeek V4-style KV compression: compress every m tokens into 1 KV entry
# via learned weighted pooling (W_KV, W_Z, positional bias B).
# Combined with MQA (kv=1) and FP8 storage, reduces KV cache by m * 2 * 2.
kv_compress_m: int = 4 # compression rate (1 = no compression, 4 = 4 tokens → 1 entry)
kv_fp8: bool = True # store compressed KV in FP8 (E4M3) instead of bf16
sparse_stride: int = 4 # within-block sparse: attend to every Nth compressed entry (1 = dense)
raw_window: int = 4096 # exact raw-KV sliding window (recent ctx, uncompressed);
# 0=compressed-only, 2048=115K tps, 4096=~99K, 8192=~76K @20M
@property
def head_dim(self) -> int:
return self.d_model // self.n_heads
@property
def medusa_par_heads(self) -> int:
"""Parallel tail heads = total - conditioned."""
return self.medusa_heads - self.medusa_cond_heads
@property
def medusa_cond_in(self) -> int:
"""Conditioned-head input width: h_t plus G compressed token embs."""
return self.d_model + self.medusa_cond_group * self.medusa_emb_rank
@property
def medusa_n_groups(self) -> int:
return self.medusa_cond_heads // self.medusa_cond_group
@property
def intermediate(self) -> int:
return int(round(self.ffn_mult * self.d_model))
@property
def block_size_comp(self) -> int:
"""Compressed entries per MoBA block = (K+1) // m.
Requires K+1 divisible by m for clean 1-block-per-step alignment."""
return (self.medusa_heads + 1) // self.kv_compress_m
@property
def kv_cache_dtype_size(self) -> int:
"""Bytes per KV element: 1 for FP8, 2 for bf16."""
return 1 if self.kv_fp8 else 2
@classmethod
def v5_500m(cls) -> "Config":
"""~436M total params: 272M base + 164M hybrid spec heads.
v6 heads (replaces rank-1 medusa):
- 256 conditioned heads (32 groups x 8): head k predicts offset k+1
given h_t + compressed embs of previous group's tokens.
- 256 parallel tail heads (offsets 257..512) from h_t alone.
- Shared vocab projection [256, 50257] + bias + token compressor.
Draft params: 83.9M cond + 67.1M par + 12.9M vocab + ~0.1M proj.
"""
return cls(
d_model=1024, n_layers=8, n_heads=16, n_kv_heads=1,
ffn_mult=8.0, medusa_heads=512, medusa_cond_heads=256,
medusa_cond_group=8, medusa_rank=256, medusa_emb_rank=32,
medusa_hidden=0, medusa_hi_heads=8, medusa_hi_rank=1024,
medusa_seq_len=32,
window_size=0, max_seq_len=20_000_000 + 5000,
use_alibi=False, dtype="bfloat16",
moba_block_size=4096, moba_top_k=3,
kv_compress_m=4, kv_fp8=True, sparse_stride=4,
raw_window=4096,
)
def estimate_vram_gb(self, batch: int = 1, gen_tokens: int = 512) -> float:
"""Rough VRAM estimate in GB before allocating — prevents hard OOM crashes.
Components:
- params (bf16, 2 bytes): base + medusa heads
- KV cache (bf16): batch * max_seq * n_kv * hd * 2(K+V) * 2 bytes * n_layers
- per-step activations (bf16): the (K+1)-token forward pass through n_layers.
Dominated by FFN intermediates (gate/up/down each [T, inter]).
"""
p = self.count_params()
params_gb = p["grand_total"] * 2 / 1e9
# KV cache uses window_size (sliding) if set, else max_seq_len (growing)
cache_len = self.window_size if self.window_size > 0 else self.max_seq_len
kv = (batch * cache_len * self.n_kv_heads * self.head_dim
* 2 * 2 * self.n_layers) / 1e9
T = self.medusa_heads + 1
# FFN activations: 3 matrices of [T, inter] per layer, plus attention I/O
actn = (3 * T * self.intermediate + 2 * T * self.d_model) * self.n_layers * 2 / 1e9
# attention QKVO projections
actn += (4 * T * self.d_model * self.n_layers + T * self.d_model) * 2 / 1e9
return params_gb + kv + actn
def count_params(self) -> dict:
d, L = self.d_model, self.n_layers
v, inter = self.vocab_size, self.intermediate
h, kv, hd = self.n_heads, self.n_kv_heads, self.head_dim
emb = v * d # tied: counts once
# attention per layer: q (d->h*hd), k/v (d->kv*hd), o (h*hd->d)
attn = d * (h * hd) + 2 * (d * (kv * hd)) + (h * hd) * d
# SwiGLU per layer: gate (d->inter), up (d->inter), down (inter->d)
ffn = 3 * d * inter
# RMSNorm params (2 per layer + final)
norm = 2 * d
base = emb + L * (attn + ffn + norm)
# Spec heads: two modes.
# r == 1 (legacy): per-head down [K,d,1] + out [K,1,v] — 2-token argmax trick.
# r > 1 (v6 hybrid): conditioned heads [Kc, d+G*e, r] + parallel [Kp,d,r]
# + shared vocab [r,v] + bias [v] + token compressor [d,e].
r, hid = self.medusa_rank, self.medusa_hidden
if r == 1:
per_head = d * r + (r * r if hid > 0 else 0) + r * v
medusa = self.medusa_heads * per_head
else:
e, g = self.medusa_emb_rank, self.medusa_cond_group
kc, kp = self.medusa_cond_heads, self.medusa_par_heads
cond = kc * (d + g * e) * r + (kc * r * r if hid > 0 else 0)
par = kp * d * r
shared = r * v + v + d * e # vocab proj + bias + tok compressor
medusa = cond + par + shared
per_head = medusa // max(1, self.medusa_heads)
return {
"embedding": emb,
"per_layer_attn": attn,
"per_layer_ffn": ffn,
"base_total": base,
"medusa_per_head": per_head,
"medusa_total": medusa,
"grand_total": base + medusa,
"tokens_per_step": self.medusa_heads + 1,
}
if __name__ == "__main__":
c = Config()
p = c.count_params()
print(f"Config: d={c.d_model} L={c.n_layers} heads={c.n_heads} kv={c.n_kv_heads} "
f"inter={c.intermediate} vocab={c.vocab_size}")
print(f" base params : {p['base_total']/1e6:7.2f}M")
print(f" medusa heads: {c.medusa_heads} x {p['medusa_per_head']/1e6:.3f}M = {p['medusa_total']/1e6:6.2f}M")
print(f" GRAND TOTAL : {p['grand_total']/1e6:7.2f}M")
print(f" tokens/step : {p['tokens_per_step']} (accept-all speculative)")