File size: 10,831 Bytes
a8f07a3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 | """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)")
|