File size: 5,486 Bytes
fdc6474
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Byte-exact VRAM budget for GLM-5.2 hybrid quant on 4x RTX PRO 6000 (96GB, SM120).

Model constants from zai-org/GLM-5.2 config.json (verified against the
lukealonso/GLM-5.2-NVFP4 safetensors index: predicted 450.99 GB vs actual
450.93 GB total_size).
"""
GB = 1e9
GiB = 2**30

# --- architecture ---
H = 6144                 # hidden_size
MOE_I = 2048             # moe_intermediate_size
DENSE_I = 12288          # intermediate_size (dense layers 0-2)
VOCAB = 154880
N_LAYERS = 78            # main layers
N_MTP = 1                # layer 78
N_DENSE = 3              # layers 0,1,2
N_EXPERT_LAYERS = 75 + 1 # layers 3..77 + MTP layer
N_EXPERTS = 256
KV_LORA = 512
QK_ROPE = 64
Q_LORA = 2048
N_HEADS = 64
QK_NOPE = 192
V_HEAD = 256
IDX_HEADS, IDX_HDIM = 32, 128
N_FULL_INDEXERS = 22     # 21 main (IndexShare freq=4) + 1 MTP
CTX = 1_048_576

# --- per-component parameter counts ---
expert_params_down = MOE_I * H            # 12.58M
expert_params_gateup = 2 * H * MOE_I      # 25.17M
expert_params = expert_params_down + expert_params_gateup   # 37.75M
layer_expert_params = N_EXPERTS * expert_params             # 9.664B

attn_params = (H*Q_LORA + Q_LORA*(N_HEADS*(QK_NOPE+QK_ROPE))
               + H*(KV_LORA+QK_ROPE) + KV_LORA*(N_HEADS*(QK_NOPE+V_HEAD))
               + (N_HEADS*V_HEAD)*H)      # 165.03M per layer
indexer_params = Q_LORA*IDX_HEADS*IDX_HDIM + H*IDX_HDIM + H*IDX_HEADS  # 9.37M
shared_expert_params = expert_params
dense_mlp_params = 3 * H * DENSE_I
gate_params = H * N_EXPERTS
embed_params = VOCAB * H

TOTAL_LAYER_ENTRIES = N_LAYERS + N_MTP    # 79

nonexpert = {
    'attention':      attn_params * TOTAL_LAYER_ENTRIES,
    'indexers':       indexer_params * N_FULL_INDEXERS,
    'dense_mlp':      dense_mlp_params * N_DENSE,
    'shared_experts': shared_expert_params * N_EXPERT_LAYERS,
    'gates_fp32':     gate_params * N_EXPERT_LAYERS,      # router kept fp32
    'embed+lm_head':  2 * embed_params,
    'mtp_extras':     2 * H * H + 4 * H,                  # eh_proj etc
}

# --- bytes/param for each quant recipe (routed experts) ---
def gguf_bpw(down_bpw, gateup_bpw):
    return (expert_params_down*down_bpw + expert_params_gateup*gateup_bpw) / expert_params

RECIPES = {
    # NVFP4: 4b weight + fp8 scale per 16 = 4.5 bpw (+ negligible per-tensor scales)
    'nvfp4':  4.5,
    # unsloth UD-Q3_K_XL donor: down=IQ4_XS(4.25), gate/up=IQ3_XXS(3.0625)
    'ud_q3':  gguf_bpw(4.25, 3.0625),
    # unsloth UD-Q2_K_XL donor: down=IQ3_XXS(3.0625), gate/up=IQ2_XS(2.3125)
    'ud_q2':  gguf_bpw(3.0625, 2.3125),
}

def layer_expert_bytes(recipe):
    return layer_expert_params * RECIPES[recipe] / 8

# --- KV cache at 1M ---
kv_latent_per_tok = (KV_LORA + QK_ROPE) * TOTAL_LAYER_ENTRIES        # values (MTP has KV too)
def kv_bytes(bytes_per_val, scale_overhead=0.02):
    return kv_latent_per_tok * bytes_per_val * CTX * (1 + scale_overhead)
idx_cache = N_FULL_INDEXERS * IDX_HDIM * 1 * CTX   # fp8 indexer k-cache

# --- GPU budget ---
GPU_MEM = 97887 * 2**20        # RTX PRO 6000 Blackwell, nvidia-smi reported
N_GPU = 4
RESERVE_PER_GPU = 6 * GiB      # cuda ctx + nccl + cudagraphs + activations + frag

def report(attn_dtype_bytes=2, n4=0, n3=0, n2=0, kv='fp8'):
    assert n4+n3+n2 == 75
    ne = dict(nonexpert)
    scale = {'attention': attn_dtype_bytes/2, 'shared_experts': attn_dtype_bytes/2,
             'dense_mlp': attn_dtype_bytes/2}
    ne_bytes = 0
    for k, p in ne.items():
        b = p * (4 if k=='gates_fp32' else 2)
        b *= scale.get(k, 1)
        ne_bytes += b
    expert_bytes = (n4*layer_expert_bytes('nvfp4') + n3*layer_expert_bytes('ud_q3')
                    + n2*layer_expert_bytes('ud_q2') + layer_expert_bytes('nvfp4'))  # +MTP nvfp4
    weights = ne_bytes + expert_bytes
    kvb = kv_bytes(1 if kv=='fp8' else 0.5) + idx_cache
    total_need = weights + kvb
    total_have = N_GPU * (GPU_MEM - RESERVE_PER_GPU)
    print(f'--- attn={16*attn_dtype_bytes//2}bit  experts: {n4}xNVFP4 {n3}xUD-Q3 {n2}xUD-Q2 (+MTP nvfp4), KV={kv} ---')
    print(f'  non-expert weights: {ne_bytes/GB:8.2f} GB')
    print(f'  expert weights:     {expert_bytes/GB:8.2f} GB   (avg {expert_bytes*8/ (N_EXPERT_LAYERS*layer_expert_params):.2f} bpw)')
    print(f'  KV latent @1M:      {kv_bytes(1 if kv=="fp8" else 0.5)/GB:8.2f} GB')
    print(f'  indexer cache @1M:  {idx_cache/GB:8.2f} GB')
    print(f'  TOTAL NEED:         {total_need/GB:8.2f} GB = {total_need/GiB:.1f} GiB')
    print(f'  TOTAL HAVE (4x{GPU_MEM/GiB:.1f}GiB - 4x{RESERVE_PER_GPU/GiB:.0f}GiB): {total_have/GB:.2f} GB = {total_have/GiB:.1f} GiB')
    margin = total_have - total_need
    print(f'  MARGIN:             {margin/GB:+8.2f} GB   {"FITS" if margin>0 else "DOES NOT FIT"}')
    # PP=4 per-stage worst case: stage holding embed+lm_head + its layers
    return margin

print('kv latent values/token:', kv_latent_per_tok, '->', kv_latent_per_tok, 'B/tok fp8,',
      f'{kv_bytes(1)/GB:.1f} GB @1M fp8')
print('nonexpert param count:', sum(nonexpert.values())/1e9, 'B params')
print('checkpoint check: nvfp4-all =',
      (sum(v*2 for v in nonexpert.values()) + 2*nonexpert['gates_fp32'] +
       76*layer_expert_bytes('nvfp4'))/GB, 'GB vs actual 450.93')
print()
print('=== v1: BF16 attention (as shipped in NVFP4 repo) ===')
report(2, 10, 30, 35)
report(2, 8, 26, 41)
report(2, 6, 22, 47)
print()
print('=== v2: FP8 attention/shared/dense (unsloth uses Q8_0 there too) ===')
report(1, 14, 40, 21)
report(1, 12, 36, 27)
report(1, 10, 30, 35)