File size: 3,183 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 | #!/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")
|