christopher-kapic's picture
Upload folder using huggingface_hub
fdc6474 verified
Raw
History Blame Contribute Delete
8.5 kB
#!/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)