#!/usr/bin/env python3 """Kernel-level golden vectors for the aqlm_moe extension, using REAL weights. For layers 6/40/76 of /data/glm52: loads the layer's hybrid tensors, runs all four kernel entry points (aqlm_moe_gemv, aqlm_moe_dequant, nvfp4_moe_gemv, nvfp4_moe_dequant) on fixed-seed fp16 inputs, and stores inputs + outputs + pure-torch reference outputs. An SM120 build of the extension must reproduce these within tolerance (see verify_sm120.py). """ import json import os import sys import torch sys.path.insert(0, "/home/coder/git/glm52/vllm") GOLD = "/data/glm52-sm120-golden" CKPT = "/data/glm52" LAYERS = [6, 40, 76] N_TOKENS = 24 def load_layer(li, device): 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", "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 main(): from vllm.model_executor.layers.quantization.nvfp4_aqlm_hybrid import ( _dequant_reference, _get_ext, ) device = "cuda:0" ext = _get_ext() torch.manual_seed(20260706) out = {} for li in LAYERS: t = load_layer(li, device) kind = t["hyb_kind"] nb = int((kind == 2).sum()) na = int((kind == 0).sum()) h = 6144 i = 2048 x_h = (torch.randn(N_TOKENS, h, dtype=torch.float16, device=device) * 0.05) x_i = (torch.randn(N_TOKENS, i, dtype=torch.float16, device=device) * 0.05) aq_ids = torch.randint(0, nb, (N_TOKENS,), dtype=torch.int32, device=device) aq_ids[0] = -1 nv_ids = torch.randint(0, na, (N_TOKENS,), dtype=torch.int32, device=device) nv_ids[1] = -1 deq_aq = torch.arange(0, min(nb, 4), dtype=torch.int32, device=device) deq_nv = torch.arange(0, min(na, 4), dtype=torch.int32, device=device) rec = {"layer": li, "n_nvfp4": na, "n_cold": nb, "x_h": x_h.cpu(), "x_i": x_i.cpu(), "aq_ids": aq_ids.cpu(), "nv_ids": nv_ids.cpu(), "deq_aq": deq_aq.cpu(), "deq_nv": deq_nv.cpu()} rec["aqlm_gemv_w13"] = ext.aqlm_moe_gemv( x_h, t["w13_codes"], t["w13_codebooks"], t["w13_scales"], aq_ids ).cpu() rec["aqlm_gemv_w2c"] = ext.aqlm_moe_gemv( x_i, t["w2c_codes"], t["w2c_codebooks"], t["w2c_scales"], aq_ids ).cpu() rec["aqlm_dequant_w13"] = ext.aqlm_moe_dequant( t["w13_codes"], t["w13_codebooks"], t["w13_scales"], deq_aq ).cpu() rec["nvfp4_gemv_w13"] = ext.nvfp4_moe_gemv( x_h, t["nvfp4_w13_packed"], t["nvfp4_w13_bscale"], t["nvfp4_w13_scale2"], nv_ids).cpu() rec["nvfp4_gemv_w2"] = ext.nvfp4_moe_gemv( x_i, t["nvfp4_w2_packed"], t["nvfp4_w2_bscale"], t["nvfp4_w2_scale2"], nv_ids).cpu() rec["nvfp4_dequant_w2"] = ext.nvfp4_moe_dequant( t["nvfp4_w2_packed"], t["nvfp4_w2_bscale"], t["nvfp4_w2_scale2"], deq_nv).cpu() # pure-torch reference for the AQLM dequant (arch-independent truth) rec["aqlm_dequant_w13_ref"] = _dequant_reference( t["w13_codes"][deq_aq.long()], t["w13_codebooks"], t["w13_scales"][deq_aq.long()]).half().cpu() out[li] = rec print(f"layer {li}: vectors done (na={na}, nb={nb})") os.makedirs(GOLD, exist_ok=True) torch.save(out, os.path.join(GOLD, "kernel_vectors.pt")) print("saved", os.path.join(GOLD, "kernel_vectors.pt")) if __name__ == "__main__": main()