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