| |
| """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 ( |
| _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" |
| |
| 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) |
| |
| 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") |
|
|