"""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)")