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