Download code/config.py from Akahsizrr/spec100m: direct link, hf CLI and curl.
- Browser
- Download file 10.8 kB
-
https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/config.py
- Command line
-
hf download hf://Akahsizrr/spec100m/code/config.py
-
curl -L -o config.py https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/config.py
10.8 kB
| """Model configuration + parameter counter for the ~100M speculative-decoding model.""" | |
| from dataclasses import dataclass, field | |
| 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 | |
| def head_dim(self) -> int: | |
| return self.d_model // self.n_heads | |
| def medusa_par_heads(self) -> int: | |
| """Parallel tail heads = total - conditioned.""" | |
| return self.medusa_heads - self.medusa_cond_heads | |
| 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 | |
| def medusa_n_groups(self) -> int: | |
| return self.medusa_cond_heads // self.medusa_cond_group | |
| def intermediate(self) -> int: | |
| return int(round(self.ffn_mult * self.d_model)) | |
| 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 | |
| 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 | |
| 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)") | |