#!/usr/bin/env python3 """ Parameter / KV-cache / attention-FLOP calculator for * DeepSeek-V3-style configs (Moonlight-16B-A3B, DeepSeek-V3) * DeepSeek-V4-style configs (DeepSeek-V4-Flash / -Pro, Moonlight-V4) Formulas follow the tensor shapes of the released checkpoints: * Moonlight: modeling_deepseek.py (DeepseekV3) shapes read from model.safetensors headers * DeepSeek-V4: inference/model.py (ModelArgs / Block / Attention / Compressor / Indexer / MoE) Published numbers used for validation (all reproduced by this script): Moonlight-16B-A3B : 15.29B total / 2.24B activated non-embedding; 16B / 3B incl. embedding DeepSeek-V4-Flash : 284B total / 13B activated DeepSeek-V4-Pro : 1.6T total / 49B activated Usage: python3 count_params.py # validation + Moonlight-V4 report python3 count_params.py path/to/config.json [...] # report for arbitrary HF-style configs """ import json import math import random import sys from pathlib import Path HERE = Path(__file__).resolve().parent # ---------------------------------------------------------------------------------------------- # DeepSeek-V3 architecture (Moonlight, DeepSeek-V3) # ---------------------------------------------------------------------------------------------- def count_v3(c): V, D, L = c["vocab_size"], c["hidden_size"], c["num_hidden_layers"] H = c["num_attention_heads"] qk_nope, qk_rope, v_hd = c["qk_nope_head_dim"], c["qk_rope_head_dim"], c["v_head_dim"] kv_r, q_r = c["kv_lora_rank"], c.get("q_lora_rank") E, K, I = c["n_routed_experts"], c["num_experts_per_tok"], c["moe_intermediate_size"] n_sh, I_dense = c["n_shared_experts"], c["intermediate_size"] k_dense, moe_freq = c["first_k_dense_replace"], c.get("moe_layer_freq", 1) n_mtp = c.get("num_nextn_predict_layers", 0) # attention (MLA) if q_r: attn = D * q_r + q_r + q_r * H * (qk_nope + qk_rope) # q_a_proj, q_a_layernorm, q_b_proj else: attn = D * H * (qk_nope + qk_rope) # q_proj attn += D * (kv_r + qk_rope) + kv_r # kv_a_proj_with_mqa, kv_a_layernorm attn += kv_r * H * (qk_nope + v_hd) # kv_b_proj attn += H * v_hd * D # o_proj norms = 2 * D dense_mlp = 3 * D * I_dense gate = E * D + E # gate.weight + e_score_correction_bias expert = 3 * D * I moe_total = gate + E * expert + n_sh * expert moe_active = gate + K * expert + n_sh * expert def is_moe(l): return l >= k_dense and (l - k_dense) % moe_freq == 0 if moe_freq > 1 else l >= k_dense n_moe = sum(1 for l in range(L) if is_moe(l)) n_dense = L - n_moe body_total = L * (attn + norms) + n_dense * dense_mlp + n_moe * moe_total + D body_active = L * (attn + norms) + n_dense * dense_mlp + n_moe * moe_active + D emb = V * D head = V * D # MTP module (DeepSeek-V3 style): eh_proj (2D x D) + enorm + hnorm + shared_head.norm + one full layer mtp_total = n_mtp * (2 * D * D + 3 * D + attn + norms + moe_total) mtp_active = n_mtp * (2 * D * D + 3 * D + attn + norms + moe_active) return dict( family="deepseek_v3", layers=L, n_dense=n_dense, n_moe=n_moe, attn_per_layer=attn, dense_mlp=dense_mlp, moe_total_per_layer=moe_total, moe_active_per_layer=moe_active, expert=expert, embed=emb, head=head, body_total=body_total, body_active=body_active, total_nonembed=body_total, active_nonembed=body_active, total_incl_embed=body_total + emb + head, active_incl_embed=body_active + emb + head, mtp_total=mtp_total, mtp_active=mtp_active, ) # ---------------------------------------------------------------------------------------------- # DeepSeek-V4 architecture (Flash, Pro, Moonlight-V4) # ---------------------------------------------------------------------------------------------- def v4_attention_params(c, ratio): """Parameters of one Attention module (incl. compressor / indexer) for a layer with the given compress ratio.""" D, H, hd = c["hidden_size"], c["num_attention_heads"], c["head_dim"] Qr, Or, G = c["q_lora_rank"], c["o_lora_rank"], c["o_groups"] IH, Ihd = c["index_n_heads"], c["index_head_dim"] p = {} p["attn_sink"] = H p["wq_a"] = D * Qr p["q_norm"] = Qr p["wq_b"] = Qr * H * hd p["wkv"] = D * hd p["kv_norm"] = hd p["wo_a"] = (H * hd // G) * (G * Or) # == H*hd*Or, applied group-wise as [G, Or, H*hd/G] p["wo_b"] = G * Or * D if ratio: coff = 2 if ratio == 4 else 1 # overlapped compression only for the CSA ratio (4) p["compressor"] = ratio * coff * hd + 2 * D * coff * hd + hd # ape + wkv + wgate + norm if ratio == 4: # CSA: lightning indexer on top of compressed keys idx = Qr * IH * Ihd + D * IH # wq_b + weights_proj idx += 4 * 2 * Ihd + 2 * D * 2 * Ihd + Ihd # indexer.compressor (ratio 4, overlap) p["indexer"] = idx return p def v4_hc_params(c): hc, D = c["hc_mult"], c["hidden_size"] mix = (2 + hc) * hc return 2 * (mix * hc * D + mix + 3) # (fn + base + scale) for attn and for ffn def count_v4(c): V, D, L = c["vocab_size"], c["hidden_size"], c["num_hidden_layers"] E, K, I = c["n_routed_experts"], c["num_experts_per_tok"], c["moe_intermediate_size"] n_sh, n_hash = c["n_shared_experts"], c["num_hash_layers"] hc = c["hc_mult"] n_mtp = c.get("num_nextn_predict_layers", 0) ratios = list(c["compress_ratios"]) assert len(ratios) >= L + n_mtp, f"compress_ratios needs >= {L + n_mtp} entries, got {len(ratios)}" expert = 3 * D * I gate_w, gate_b = E * D, E moe_total = lambda hash_layer: gate_w + (0 if hash_layer else gate_b) + E * expert + n_sh * expert moe_active = lambda hash_layer: gate_w + (0 if hash_layer else gate_b) + K * expert + n_sh * expert hc_per_layer = v4_hc_params(c) norms = 2 * D per_layer = [] for l in range(L): ap = v4_attention_params(c, ratios[l]) attn = sum(ap.values()) is_hash = l < n_hash per_layer.append(dict( layer=l, ratio=ratios[l], hash=is_hash, attn=attn, attn_parts=ap, hc=hc_per_layer, total=attn + hc_per_layer + norms + moe_total(is_hash), active=attn + hc_per_layer + norms + moe_active(is_hash), )) head_hc = hc * hc * D + hc + 1 body_total = sum(x["total"] for x in per_layer) + D + head_hc body_active = sum(x["active"] for x in per_layer) + D + head_hc emb, head = V * D, V * D mtp_total = mtp_active = 0 for i in range(n_mtp): ap = sum(v4_attention_params(c, ratios[L + i]).values()) extra = 2 * D * D + 3 * D + head_hc # e_proj, h_proj, enorm, hnorm, norm, hc_head_* mtp_total += ap + hc_per_layer + norms + moe_total(False) + extra mtp_active += ap + hc_per_layer + norms + moe_active(False) + extra kinds = {0: "SWA", 4: "CSA", 128: "HCA"} layer_kinds = {k: sum(1 for l in range(L) if kinds.get(ratios[l]) == k) for k in kinds.values()} return dict( family="deepseek_v4", layers=L, layer_kinds=layer_kinds, n_hash=n_hash, per_layer=per_layer, expert=expert, hc_per_layer=hc_per_layer, moe_total_per_layer=moe_total(False), moe_active_per_layer=moe_active(False), embed=emb, head=head, body_total=body_total, body_active=body_active, total_nonembed=body_total, active_nonembed=body_active, total_incl_embed=body_total + emb + head, active_incl_embed=body_active + emb + head, mtp_total=mtp_total, mtp_active=mtp_active, tid2eid_entries=n_hash * V * K, # non-trainable hash-routing lookup tables ) def count(c): return count_v4(c) if c.get("model_type") == "deepseek_v4" or "compress_ratios" in c else count_v3(c) # ---------------------------------------------------------------------------------------------- # KV cache and attention FLOPs per token as a function of context length # ---------------------------------------------------------------------------------------------- def kv_cache_bytes_v3(c, L, kv_bytes=2): """MLA cache: one (kv_lora_rank + rope) latent per token per layer.""" per_tok = (c["kv_lora_rank"] + c["qk_rope_head_dim"]) * kv_bytes return c["num_hidden_layers"] * per_tok * L def kv_cache_bytes_v4(c, L, mixed=True): """V4 cache. Per entry: head_dim dims; mixed storage = fp8 for non-rope dims + bf16 for the 64 rope dims (paper sec. 2.3.4). CSA layers additionally cache indexer keys (index_head_dim, FP4).""" hd, rd = c["head_dim"], c["qk_rope_head_dim"] win = c["sliding_window"] entry = (hd - rd) * 1 + rd * 2 if mixed else hd * 2 idx_entry = c["index_head_dim"] * (0.5 if mixed else 2) total = 0 for r in c["compress_ratios"][: c["num_hidden_layers"]]: n = min(win, L) if r: n += L // r total += n * entry if r == 4: total += (L // r) * idx_entry return total def attn_flops_v3(c, L): """Core attention FLOPs for ONE query token attending to L cached tokens, all layers. Naive (non-absorbed) MLA: QK over (nope+rope) dims and PV over v dims, per head.""" H = c["num_attention_heads"] qk = c["qk_nope_head_dim"] + c["qk_rope_head_dim"] return c["num_hidden_layers"] * 2 * H * (qk + c["v_head_dim"]) * L def attn_flops_v4(c, L): """Core attention (+ lightning indexer) FLOPs for ONE query token at context L, all layers. KV entries serve as both key and value: 2*H*hd per entry for QK and 2*H*hd for PV.""" H, hd, win = c["num_attention_heads"], c["head_dim"], c["sliding_window"] IH, Ihd, topk = c["index_n_heads"], c["index_head_dim"], c["index_topk"] total = 0 for r in c["compress_ratios"][: c["num_hidden_layers"]]: n = min(win, L) if r == 4: n += min(L // r, topk) total += 2 * IH * Ihd * (L // r) # indexer scores over all compressed keys elif r: n += L // r total += 4 * H * hd * n return total # ---------------------------------------------------------------------------------------------- # Routed scaling factor, Moonlight's recipe (appendix C, fig. 6) generalised to any scoring function # ---------------------------------------------------------------------------------------------- def gate_scaling_factor(num_experts, topk, score="sigmoid", iters=200_000, seed=0): rng = random.Random(seed) if score == "sigmoid": f = lambda x: 1.0 / (1.0 + math.exp(-x)) elif score == "sqrtsoftplus": f = lambda x: math.sqrt(math.log1p(math.exp(x))) else: raise ValueError(score) acc = 0.0 for _ in range(iters): p = sorted((f(rng.gauss(0, 1)) for _ in range(num_experts)), reverse=True)[:topk] s = sum(p) acc += 1.0 / math.sqrt(sum((x / s) ** 2 for x in p)) return acc / iters # ---------------------------------------------------------------------------------------------- def fmt(n): if n >= 1e12: return f"{n/1e12:.3f}T" if n >= 1e9: return f"{n/1e9:.3f}B" if n >= 1e6: return f"{n/1e6:.2f}M" return f"{n:,}" def gb(x): return f"{x/2**30:.3f} GiB" if x >= 2**30 else f"{x/2**20:.1f} MiB" def report(name, c): r = count(c) print(f"== {name} ({r['family']}) ==") print(f" layers={r['layers']}" + (f" kinds={r['layer_kinds']} hash_layers={r['n_hash']}" if r["family"] == "deepseek_v4" else f" dense={r['n_dense']} moe={r['n_moe']}")) print(f" embed={fmt(r['embed'])} head={fmt(r['head'])} expert={fmt(r['expert'])}") print(f" total non-embedding: {fmt(r['total_nonembed'])} incl. embed+head: {fmt(r['total_incl_embed'])}") print(f" active non-embedding: {fmt(r['active_nonembed'])} incl. embed+head: {fmt(r['active_incl_embed'])}") if r["mtp_total"]: print(f" MTP module(s): total {fmt(r['mtp_total'])}, active {fmt(r['mtp_active'])}" f" -> grand total incl. MTP {fmt(r['total_incl_embed'] + r['mtp_total'])}") if r["family"] == "deepseek_v4": for kind in ("SWA", "CSA", "HCA"): ex = next((x for x in r["per_layer"] if {0: "SWA", 4: "CSA", 128: "HCA"}[x["ratio"]] == kind), None) if ex: parts = ", ".join(f"{k}={fmt(v)}" for k, v in ex["attn_parts"].items()) print(f" attention/{kind} layer: {fmt(ex['attn'])} [{parts}]") print(f" mHC per layer: {fmt(r['hc_per_layer'])} tid2eid entries (non-trainable): {fmt(r['tid2eid_entries'])}") else: print(f" attention per layer: {fmt(r['attn_per_layer'])} dense MLP: {fmt(r['dense_mlp'])}") print(f" MoE per layer: total {fmt(r['moe_total_per_layer'])}, active {fmt(r['moe_active_per_layer'])}") return r def main(argv): if argv: for p in argv: report(Path(p).stem, json.load(open(p))) return ref = HERE / "reference" moonlight = json.load(open(ref / "moonlight_16b_a3b_config.json")) flash = json.load(open(ref / "deepseek_v4_flash_config.json")) pro = json.load(open(ref / "deepseek_v4_pro_config.json")) mv4 = json.load(open(HERE / "config.json")) print("#" * 100 + "\n# Validation against published numbers\n" + "#" * 100) r_ml = report("Moonlight-16B-A3B (published: 15.29B/2.24B non-embed, 16B/3B incl. embed)", moonlight) r_fl = report("DeepSeek-V4-Flash (published: 284B total, 13B active)", flash) r_pr = report("DeepSeek-V4-Pro (published: 1.6T total, 49B active)", pro) print("\n" + "#" * 100 + "\n# Moonlight-V4 (this repo)\n" + "#" * 100) r_m4 = report("Moonlight-V4-16B-A3B", mv4) print("\n#### Routed scaling factor via Moonlight's recipe (E[1/||p||_2] over renormalised top-k scores)") for (E, K, s) in [(64, 6, "sigmoid"), (64, 6, "sqrtsoftplus"), (256, 6, "sqrtsoftplus"), (384, 6, "sqrtsoftplus"), (256, 8, "sigmoid")]: print(f" experts={E:4d} topk={K} score={s:12s}: {gate_scaling_factor(E, K, s, iters=20000):.3f}") print("\n#### KV cache per sequence (Moonlight: bf16 MLA latent; V4: fp8 non-rope + bf16 rope + fp4 indexer keys)") for L in (8192, 65536, 1048576): print(f" L={L:>8}: Moonlight {gb(kv_cache_bytes_v3(moonlight, L)):>12} | Moonlight-V4 {gb(kv_cache_bytes_v4(mv4, L)):>12}" f" (bf16-only {gb(kv_cache_bytes_v4(mv4, L, mixed=False)):>12}) | V4-Flash {gb(kv_cache_bytes_v4(flash, L)):>12}") print("\n#### Per-token FLOPs at context L (2*active params incl. head + core attention [+ indexer])") for L in (8192, 65536, 1048576): lin_ml = 2 * (r_ml["active_nonembed"] + r_ml["head"]) lin_m4 = 2 * (r_m4["active_nonembed"] + r_m4["head"]) a_ml, a_m4 = attn_flops_v3(moonlight, L), attn_flops_v4(mv4, L) print(f" L={L:>8}: Moonlight linear {lin_ml/1e9:6.2f} GF + attn {a_ml/1e9:8.2f} GF = {(lin_ml+a_ml)/1e9:8.2f} GF" f" | Moonlight-V4 linear {lin_m4/1e9:6.2f} GF + attn {a_m4/1e9:6.2f} GF = {(lin_m4+a_m4)/1e9:6.2f} GF") if __name__ == "__main__": main(sys.argv[1:])