#!/usr/bin/env python3 """SC-2: dequant statistics on real weights (1 GPU, ~2 min). Usage: sc2_dequant_stats.py CKPT [--layers 3,21,40,60,77] sc2_dequant_stats.py --parts PARTS_DIR [--layers ...]""" import argparse, json, sys import torch sys.path.insert(0, "/home/coder/git/glm52/vllm") from vllm.model_executor.layers.quantization.nvfp4_aqlm_hybrid import ( # noqa _dequant_reference) FP4_LUT = torch.tensor([0.0,0.5,1.0,1.5,2.0,3.0,4.0,6.0, -0.0,-0.5,-1.0,-1.5,-2.0,-3.0,-4.0,-6.0]) ap = argparse.ArgumentParser() ap.add_argument("ckpt", nargs="?") ap.add_argument("--parts") ap.add_argument("--layers", default="3,21,40,60,77") a = ap.parse_args() layers = [int(x) for x in a.layers.split(",")] dev = "cuda:0" if torch.cuda.is_available() else "cpu" lut = FP4_LUT.to(dev) problems = [] def stats_ok(name, t): t = t.float() rms = t.pow(2).mean().sqrt().item() zf = (t == 0).float().mean().item() if not torch.isfinite(t).all(): problems.append(f"{name}: NaN/Inf") if not (1e-3 <= rms <= 1.0): problems.append(f"{name}: rms {rms:.2e}") if zf > 0.30: problems.append(f"{name}: zero-frac {zf:.2f}") return rms if a.parts: for li in layers: d = torch.load(f"{a.parts}/layer_{li}.pt", map_location=dev, weights_only=True) for j in range(0, min(4, d["w13_codes"].shape[0])): w = _dequant_reference(d["w13_codes"][j:j+1].to(dev), d["w13_codebooks"].to(dev), d["w13_scales"][j:j+1].to(dev)) stats_ok(f"parts L{li} w13[{j}]", w) print(f"L{li}: parts dequant ok") else: from safetensors import safe_open idx = json.load(open(f"{a.ckpt}/model.safetensors.index.json")) wm = idx["weight_map"]; opened = {} def get(n): s = wm[n] if s not in opened: opened[s] = safe_open(f"{a.ckpt}/{s}", framework="pt") return opened[s].get_tensor(n) for li in layers: p = f"model.layers.{li}.mlp.experts" # AQLM cold experts vs pure-torch reference codes = get(f"{p}.w13_codes")[:4].to(dev) cb = get(f"{p}.w13_codebooks").to(dev) sc = get(f"{p}.w13_scales")[:4].to(dev) w = _dequant_reference(codes, cb, sc) stats_ok(f"L{li} aqlm w13", w) # NVFP4 hot experts pk = get(f"{p}.nvfp4_w13_packed")[:4].to(dev) bs = get(f"{p}.nvfp4_w13_bscale")[:4].to(dev) s2 = get(f"{p}.nvfp4_w13_scale2")[:4].to(dev).float() lo = lut[(pk & 0xF).long()]; hi = lut[(pk >> 4).long()] vals = torch.stack([lo, hi], -1).reshape(4, pk.shape[1], -1) scale = bs.view(torch.float8_e4m3fn).float().repeat_interleave(16, -1) wn = vals * scale wn[:, :2048] *= s2[:, 0, None, None] wn[:, 2048:] *= s2[:, 1, None, None] stats_ok(f"L{li} nvfp4 w13", wn) print(f"L{li}: ckpt dequant ok " f"(aqlm rms {w.float().pow(2).mean().sqrt():.3f}, " f"nvfp4 rms {wn.float().pow(2).mean().sqrt():.3f})") if problems: print("SC2 FAIL"); [print(" ", p) for p in problems]; sys.exit(1) print("SC2 PASS")