File size: 8,496 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 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 211 212 213 214 215 216 217 218 219 | #!/usr/bin/env python3
"""Verify an SM120 build against the SM100-captured golden bundle.
Usage (on the SM120 box, with the patched vLLM venv active):
python verify_sm120.py /path/to/glm52-sm120-golden [--ckpt /path/to/model]
Stages (each independent; run what's available):
1. kernels - replay kernel_vectors.pt through the JIT extension; the
custom kernels are plain CUDA and must match to ~1e-2 rel.
2. moe - replay moe_layer*.pt through HybridExpertsMoEMethod using
real layer weights from --ckpt (loads one layer per file).
3. attention - replay attn_layer*.pt through the ACTIVE sparse-MLA
backend if a standalone replay is wired up for it; else
prints the tensor contract so a new kernel can be tested
directly against q/kv_pages/topk_indices -> attn_out.
4. e2e - after `vllm serve` is up on this box, re-run the prompts
in e2e_goldens.pt (greedy) and compare generated ids and
top-50 logprob overlap per step.
fp8_ds_mla page layout (656 B/token, block=64 tokens):
[0:512] c_kv latent, fp8 e4m3 (512 dims)
[512:528] 4x fp32 scales (one per 128-dim group of c_kv)
[528:656] k_pe rope part, bf16 (64 dims)
Attention contract: out[t] = softmax(q[t] . K[topk_indices[t]] / sqrt(576))
. V[topk_indices[t]] over the 2048 selected tokens (MQA over 576-dim
latent+rope, 512-dim value = c_kv), invalid indices (-1) masked.
"""
import argparse
import glob
import os
import sys
import torch
TOLS = {"gemv": 2e-2, "dequant": 2e-3}
def rel_err(a, b):
a, b = a.float(), b.float()
return ((a - b).norm() / b.norm().clamp_min(1e-9)).item()
def check(name, got, ref, tol):
e = rel_err(got, ref)
ok = e <= tol
print(f" {'PASS' if ok else 'FAIL'} {name}: rel_err={e:.5f} (tol {tol})")
return ok
def stage_kernels(gold):
from vllm.model_executor.layers.quantization.nvfp4_aqlm_hybrid import (
_get_ext,
)
ext = _get_ext()
vecs = torch.load(os.path.join(gold, "kernel_vectors.pt"),
map_location="cpu", weights_only=True)
print("== stage 1: kernels (needs --ckpt for weights) ==")
ok = True
from verify_sm120 import _load_layer # self-import for reuse
for li, rec in vecs.items():
t = _load_layer(ARGS.ckpt, li, "cuda:0")
x_h = rec["x_h"].cuda()
x_i = rec["x_i"].cuda()
aq = rec["aq_ids"].cuda()
nv = rec["nv_ids"].cuda()
pairs = [
("aqlm_gemv_w13", ext.aqlm_moe_gemv(
x_h, t["w13_codes"], t["w13_codebooks"], t["w13_scales"], aq),
"gemv"),
("aqlm_gemv_w2c", ext.aqlm_moe_gemv(
x_i, t["w2c_codes"], t["w2c_codebooks"], t["w2c_scales"], aq),
"gemv"),
("aqlm_dequant_w13", ext.aqlm_moe_dequant(
t["w13_codes"], t["w13_codebooks"], t["w13_scales"],
rec["deq_aq"].cuda()), "dequant"),
("nvfp4_gemv_w13", ext.nvfp4_moe_gemv(
x_h, t["nvfp4_w13_packed"], t["nvfp4_w13_bscale"],
t["nvfp4_w13_scale2"], nv), "gemv"),
("nvfp4_gemv_w2", ext.nvfp4_moe_gemv(
x_i, t["nvfp4_w2_packed"], t["nvfp4_w2_bscale"],
t["nvfp4_w2_scale2"], nv), "gemv"),
("nvfp4_dequant_w2", ext.nvfp4_moe_dequant(
t["nvfp4_w2_packed"], t["nvfp4_w2_bscale"],
t["nvfp4_w2_scale2"], rec["deq_nv"].cuda()), "dequant"),
]
print(f" layer {li}:")
for name, got, kind in pairs:
ok &= check(name, got.cpu(), rec[name], TOLS[kind])
return ok
def _load_layer(ckpt, li, device):
import json
from safetensors import safe_open
idx = json.load(open(f"{ckpt}/model.safetensors.index.json"))
wm = idx["weight_map"]
p = f"model.layers.{li}.mlp.experts"
names = ["hyb_kind", "w13_codes", "w13_codebooks", "w13_scales",
"w2m_codes", "w2m_codebooks", "w2m_scales",
"w2c_codes", "w2c_codebooks", "w2c_scales",
"nvfp4_w13_packed", "nvfp4_w13_bscale", "nvfp4_w13_scale2",
"nvfp4_w2_packed", "nvfp4_w2_bscale", "nvfp4_w2_scale2"]
t, opened = {}, {}
for n in names:
shard = wm[f"{p}.{n}"]
if shard not in opened:
opened[shard] = safe_open(f"{ckpt}/{shard}", framework="pt")
t[n] = opened[shard].get_tensor(f"{p}.{n}").to(device)
return t
def stage_moe(gold):
print("== stage 2: hybrid MoE method replay ==")
from vllm.model_executor.layers.quantization.nvfp4_aqlm_hybrid import (
HybridExpertsMoEMethod,
)
ok = True
for f in sorted(glob.glob(os.path.join(gold, "moe_layer*_call*.pt"))):
rec = torch.load(f, map_location="cpu", weights_only=True)
li = rec["layer_idx"]
t = _load_layer(ARGS.ckpt, li, "cuda:0")
class L:
activation = "silu"
for k, v in t.items():
setattr(L, k, v)
class Moe:
num_experts = 256
class moe_parallel_config:
tp_size = 1
ep_size = 1
m = HybridExpertsMoEMethod.__new__(HybridExpertsMoEMethod)
m.n_nvfp4, m.n_base, m.n_cold = (rec["n_nvfp4"], rec["n_base"],
rec["n_cold"])
m.moe = Moe()
m.layer_idx = li
m._stats_dir = None
HybridExpertsMoEMethod.process_weights_after_loading(m, L)
x = rec["x"].cuda()
out = m._apply_gemv(L, x, rec["topk_weights"].cuda(),
rec["topk_ids"].cuda()).to(rec["out"].dtype)
ok &= check(os.path.basename(f), out.cpu(), rec["out"], 3e-2)
return ok
def stage_attention(gold):
print("== stage 3: sparse-MLA attention vectors ==")
files = sorted(glob.glob(os.path.join(gold, "attn_layer*_call*.pt")))
for f in files:
rec = torch.load(f, map_location="cpu", weights_only=True)
print(f" {os.path.basename(f)}: q{tuple(rec['q'].shape)} "
f"pages{tuple(rec['kv_pages'].shape)} "
f"topk{tuple(rec['topk_indices'].shape)} "
f"-> out{tuple(rec['attn_out'].shape)}")
print(" (contract in module docstring; wire your SM120 kernel's "
"replay here and compare with rel_err <= 3e-2)")
return True
def stage_e2e(gold, port):
print("== stage 4: e2e goldens vs running server ==")
import json
import urllib.request
g = torch.load(os.path.join(gold, "e2e_goldens.pt"),
map_location="cpu", weights_only=True)
ok = True
for i, case in enumerate(g["goldens"]):
body = {"model": ARGS.ckpt,
"prompt": case["prompt_token_ids"],
"max_tokens": len(case["generated_token_ids"]),
"temperature": 0.0, "logprobs": 50}
req = urllib.request.Request(
f"http://localhost:{port}/v1/completions",
data=json.dumps(body).encode(),
headers={"Content-Type": "application/json"})
with urllib.request.urlopen(req, timeout=600) as r:
resp = json.load(r)
got_text = resp["choices"][0]["text"]
match = got_text == case["generated_text"]
# token-identical is the ideal; tiny kernel diffs may flip late
# tokens, so also report first-divergence step
print(f" case {i}: {'PASS (exact)' if match else 'text differs'}")
if not match:
print(f" golden: {case['generated_text'][:60]!r}")
print(f" got: {got_text[:60]!r}")
ok = False
return ok
if __name__ == "__main__":
ap = argparse.ArgumentParser()
ap.add_argument("gold")
ap.add_argument("--ckpt", default="/data/glm52")
ap.add_argument("--port", type=int, default=0,
help="if set, run e2e stage against a live server")
ap.add_argument("--stages", default="kernels,moe,attention")
ARGS = ap.parse_args()
sys.modules["verify_sm120"] = sys.modules["__main__"]
results = {}
for s in ARGS.stages.split(","):
fn = {"kernels": stage_kernels, "moe": stage_moe,
"attention": stage_attention}.get(s)
if fn:
results[s] = fn(ARGS.gold)
if ARGS.port:
results["e2e"] = stage_e2e(ARGS.gold, ARGS.port)
print("\nsummary:", {k: "PASS" if v else "FAIL" for k, v in results.items()})
sys.exit(0 if all(results.values()) else 1)
|