christopher-kapic's picture
Upload folder using huggingface_hub
fdc6474 verified
Raw
History Blame Contribute Delete
10.4 kB
#!/usr/bin/env python3
"""REAP-style expert saliency scoring, CPU-only (runs alongside GPU jobs).
Per expert e of every MoE layer:
score_e = SUM_{t routed to e} g_{t,e} * ||f_e(x_t)||_2 * relerr_e
g router weight (topk_weights from /data/glm52-acts)
f_e(x) down(silu(gate x) * up x) computed with TEACHER (NVFP4-dequant)
weights, fp32 on CPU
relerr_e mean of w13/w2 h-weighted 2-bit reconstruction rel-errors using
the init AQLM parts (/data/glm52-aqlm-parts cover all 256
experts/layer); h = E[x^2] over the expert's routed tokens
Teachers: layers 3,4,5,8,74-77 -> /data/glm52-old-layerwise (all experts);
other layers: cold experts from /tmp/glm52-hot-dl2 regions, hot experts
from the live checkpoint's compacted nvfp4_* arrays.
Output: /data/glm52-reap-scores.npz (scores [75,256], raw saliency,
relerr, token hits, layer_ids) + per-layer spearman vs frequency counts.
Usage: score_experts_reap.py [--layers ...] [--workers 12]
"""
import argparse
import json
import os
import time
from concurrent.futures import ProcessPoolExecutor
ACTS = "/data/glm52-acts"
CKPT = "/data/glm52"
PARTS = "/data/glm52-aqlm-parts"
OLD = "/data/glm52-old-layerwise"
DL = "/tmp/glm52-hot-dl2"
OUTDIR = "/data/glm52-reap-scores"
LOCAL_TEACHER_LAYERS = {3, 4, 5, 8, 74, 75, 76, 77}
FP4_LUT = [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]
def _readers():
import torch
from safetensors import safe_open
class Regions:
def __init__(self):
self.headers = json.load(open(f"{DL}/headers.json"))
self.wm = {n: s for s, h in self.headers.items()
for n in h["header"] if n != "__metadata__"}
self.regions = {}
for shard in self.headers:
d = f"{DL}/regions/{shard}"
regs = []
if os.path.isdir(d):
for f in os.listdir(d):
p = os.path.join(d, f)
regs.append((int(f[:-4]), p, os.path.getsize(p)))
self.regions[shard] = sorted(regs)
def get(self, name):
DT = {"U8": torch.uint8, "F8_E4M3": torch.uint8,
"BF16": torch.bfloat16, "F32": torch.float32}
shard = self.wm[name]
info = self.headers[shard]["header"][name]
b, e = info["data_offsets"]
for rb, path, sz in self.regions[shard]:
if rb <= b and e <= rb + sz:
with open(path, "rb") as fh:
fh.seek(b - rb)
buf = fh.read(e - b)
return torch.frombuffer(
bytearray(buf), dtype=DT[info["dtype"]]
).reshape(info["shape"])
raise KeyError(name)
class St:
def __init__(self, root):
idx = json.load(open(f"{root}/model.safetensors.index.json"))
self.wm = idx["weight_map"]
self.root = root
self._o = {}
def get(self, name):
shard = self.wm[name]
if shard not in self._o:
self._o[shard] = safe_open(f"{self.root}/{shard}",
framework="pt")
return self._o[shard].get_tensor(name)
return Regions(), St(OLD), St(CKPT)
def _dequant_nvfp4(packed, bscale, scale2, lut):
import torch
lo = (packed & 0x0F).long()
hi = (packed >> 4).long()
vals = torch.empty(packed.shape[0], packed.shape[1] * 2,
dtype=torch.float32)
vals[:, 0::2] = lut[lo]
vals[:, 1::2] = lut[hi]
scale = bscale.view(torch.float8_e4m3fn).float().repeat_interleave(
16, dim=1)
return vals * scale * float(scale2)
def process_layer(li, threads):
import numpy as np
import torch
torch.set_num_threads(threads)
t0 = time.time()
lut = torch.tensor(FP4_LUT, dtype=torch.float32)
outp = f"{OUTDIR}/layer_{li}.npz"
if os.path.exists(outp):
return li, "skip"
regions, old, live = _readers()
acts = torch.load(f"{ACTS}/acts_layer{li}.pt", map_location="cpu",
weights_only=True)
x = acts["x"].float()
tk = acts["topk_ids"].long()
tw = acts["topk_weights"].float()
p = f"model.layers.{li}.mlp.experts"
kind = live.get(f"{p}.hyb_kind")
hot_pos = {int(e): j for j, e in
enumerate((kind == 0).nonzero().flatten().tolist())}
nv = {n: live.get(f"{p}.{n}") for n in
("nvfp4_w13_packed", "nvfp4_w13_bscale", "nvfp4_w13_scale2",
"nvfp4_w2_packed", "nvfp4_w2_bscale", "nvfp4_w2_scale2")}
part = torch.load(f"{PARTS}/layer_{li}.pt", map_location="cpu",
weights_only=True)
cb13 = part["w13_codebooks"][0].float()
cb2 = part["w2_codebooks"][0].float()
def teacher(e):
if li in LOCAL_TEACHER_LAYERS:
ep = f"model.layers.{li}.mlp.experts.{e}"
g = _dequant_nvfp4(old.get(f"{ep}.gate_proj.weight"),
old.get(f"{ep}.gate_proj.weight_scale").view(torch.uint8),
old.get(f"{ep}.gate_proj.weight_scale_2"), lut)
u = _dequant_nvfp4(old.get(f"{ep}.up_proj.weight"),
old.get(f"{ep}.up_proj.weight_scale").view(torch.uint8),
old.get(f"{ep}.up_proj.weight_scale_2"), lut)
d = _dequant_nvfp4(old.get(f"{ep}.down_proj.weight"),
old.get(f"{ep}.down_proj.weight_scale").view(torch.uint8),
old.get(f"{ep}.down_proj.weight_scale_2"), lut)
return g, u, d
if e in hot_pos:
j = hot_pos[e]
w13 = _dequant_nvfp4(nv["nvfp4_w13_packed"][j],
nv["nvfp4_w13_bscale"][j],
1.0, lut)
# scale2 is per (gate,up):
w13[:2048] *= float(nv["nvfp4_w13_scale2"][j, 0])
w13[2048:] *= float(nv["nvfp4_w13_scale2"][j, 1])
w2 = _dequant_nvfp4(nv["nvfp4_w2_packed"][j],
nv["nvfp4_w2_bscale"][j],
float(nv["nvfp4_w2_scale2"][j, 0]), lut)
return w13[:2048], w13[2048:], w2
ep = f"model.layers.{li}.mlp.experts.{e}"
g = _dequant_nvfp4(regions.get(f"{ep}.gate_proj.weight"),
regions.get(f"{ep}.gate_proj.weight_scale").view(torch.uint8),
regions.get(f"{ep}.gate_proj.weight_scale_2"), lut)
u = _dequant_nvfp4(regions.get(f"{ep}.up_proj.weight"),
regions.get(f"{ep}.up_proj.weight_scale").view(torch.uint8),
regions.get(f"{ep}.up_proj.weight_scale_2"), lut)
d = _dequant_nvfp4(regions.get(f"{ep}.down_proj.weight"),
regions.get(f"{ep}.down_proj.weight_scale").view(torch.uint8),
regions.get(f"{ep}.down_proj.weight_scale_2"), lut)
return g, u, d
def aqlm_dequant(codes, cb, scales):
idx = codes.view(torch.uint16).long() # [books,M,K8]
w = cb[idx[0]]
return w.reshape(w.shape[0], -1) * scales.unsqueeze(-1)
E = 256
sal = np.zeros(E)
rel = np.zeros(E)
hits = np.zeros(E, dtype=np.int64)
counts = np.zeros(E, dtype=np.int64)
for e in range(E):
mask = (tk == e)
rows = mask.any(1)
counts[e] = int(rows.sum())
g_w, u_w, d_w = teacher(e)
if counts[e] > 0:
xe = x[rows]
ge = tw[mask] # one weight per hit row (expert unique per token)
mid = torch.nn.functional.silu(xe @ g_w.t()) * (xe @ u_w.t())
f = mid @ d_w.t()
sal[e] = float((ge * f.norm(dim=1)).sum())
h13 = xe.pow(2).mean(0)
h2 = mid.pow(2).mean(0)
else:
h13 = x.pow(2).mean(0)
h2 = None
# relerr from init AQLM parts (per-expert, h-weighted)
w13_t = torch.cat([g_w, u_w], 0)
w13_q = aqlm_dequant(part["w13_codes"][e], cb13,
part["w13_scales"][e].float())
num = ((w13_t - w13_q).pow(2) * h13).sum()
den = (w13_t.pow(2) * h13).sum().clamp_min(1e-12)
r13 = float((num / den).sqrt())
w2_q = aqlm_dequant(part["w2_codes"][e][:1], cb2,
part["w2_scales"][e].float())
if h2 is None:
h2 = torch.ones(2048)
num2 = ((d_w - w2_q).pow(2) * h2).sum()
den2 = (d_w.pow(2) * h2).sum().clamp_min(1e-12)
r2 = float((num2 / den2).sqrt())
rel[e] = 0.5 * (r13 + r2)
hits[e] = counts[e]
score = sal * rel
from scipy.stats import spearmanr
rho = spearmanr(score, counts).statistic if counts.sum() else 0.0
np.savez(outp, score=score, saliency=sal, relerr=rel, hits=hits,
counts=counts)
return li, f"done {time.time()-t0:.0f}s spearman(score,freq)={rho:.3f}"
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--layers", default=None)
ap.add_argument("--workers", type=int, default=12)
ap.add_argument("--threads", type=int, default=9)
args = ap.parse_args()
os.makedirs(OUTDIR, exist_ok=True)
layers = list(range(3, 78))
if args.layers:
want = {int(v) for v in args.layers.split(",")}
layers = [li for li in layers if li in want]
with ProcessPoolExecutor(max_workers=args.workers) as ex:
futs = {ex.submit(process_layer, li, args.threads): li
for li in layers}
for f in futs:
pass
for f in list(futs):
li, msg = f.result()
print(f"L{li}: {msg}", flush=True)
# merge
import numpy as np
files = sorted(int(f[6:-4]) for f in os.listdir(OUTDIR)
if f.startswith("layer_"))
if len(files) == 75:
merged = {k: np.stack([np.load(f"{OUTDIR}/layer_{li}.npz")[k]
for li in files])
for k in ("score", "saliency", "relerr", "hits", "counts")}
np.savez("/data/glm52-reap-scores.npz",
layer_ids=np.array(files), **merged)
print("merged -> /data/glm52-reap-scores.npz")
if __name__ == "__main__":
main()