christopher-kapic's picture
Upload folder using huggingface_hub
fdc6474 verified
Raw
History Blame Contribute Delete
7.54 kB
#!/usr/bin/env python3
"""Phase 1.5: encode REAP-demoted experts against converged codebooks.
For each layer: experts that are HOT in the live checkpoint but COLD in
any of the three REAP assignments get AQLM codes: h-weighted encode
(2 sweeps) against the layer's CONVERGED codebook (frozen) + weighted
scale refit. Teachers = the live checkpoint's own NVFP4 hot arrays.
Output: /data/glm52-aqlm-conv15/layer_N.pt = conv parts EXTENDED with the
new experts (expert_ids re-sorted ascending).
"""
import json
import os
import time
from concurrent.futures import ProcessPoolExecutor
CKPT = "/data/glm52"
CONV = os.environ.get("P15_CONV", "/data/glm52-aqlm-conv")
OUT = os.environ.get("P15_OUT", "/data/glm52-aqlm-conv15")
ACTS = "/data/glm52-acts"
ASSIGNS = ["/data/glm52-assign-reap-250.json",
"/data/glm52-assign-reap-290.json",
"/data/glm52-assign-reap-310.json"]
G = 8
CHUNK = 32768
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 process_layer(li, gpu):
import torch
from safetensors import safe_open
torch.set_num_threads(4)
torch.cuda.set_device(gpu)
dev = f"cuda:{gpu}"
t0 = time.time()
outp = f"{OUT}/layer_{li}.pt"
if os.path.exists(outp):
return li, "skip"
lut = torch.tensor(FP4_LUT, dtype=torch.float32, device=dev)
idx = json.load(open(f"{CKPT}/model.safetensors.index.json"))
wm = idx["weight_map"]
opened = {}
def ck(name):
s = wm[f"model.layers.{li}.mlp.experts.{name}"]
if s not in opened:
opened[s] = safe_open(f"{CKPT}/{s}", framework="pt")
return opened[s].get_tensor(f"model.layers.{li}.mlp.experts.{name}")
kind = ck("hyb_kind")
cur_hot = (kind == 0).nonzero().flatten().tolist()
hot_pos = {e: j for j, e in enumerate(cur_hot)}
need_file = os.environ.get("P15_NEED")
if need_file:
need = json.load(open(need_file))
new = sorted(need.get(str(li), []))
else:
union_cold = set()
for f in ASSIGNS:
a = json.load(open(f))[str(li)]
union_cold |= set(a["cold"])
new = sorted(set(cur_hot) & union_cold) # demoted: need codes
part = torch.load(f"{CONV}/layer_{li}.pt", map_location="cpu",
weights_only=True)
if not new:
torch.save(part, outp)
return li, "no demotions, copied"
acts = torch.load(f"{ACTS}/acts_layer{li}.pt", map_location="cpu",
weights_only=True)
x = acts["x"].to(dev).float()
tk = acts["topk_ids"].to(dev).long()
nvp = {n: ck(n) for n in ("nvfp4_w13_packed", "nvfp4_w13_bscale",
"nvfp4_w13_scale2", "nvfp4_w2_packed",
"nvfp4_w2_bscale", "nvfp4_w2_scale2")}
def teacher(e):
j = hot_pos[e]
pk = nvp["nvfp4_w13_packed"][j].to(dev)
bs = nvp["nvfp4_w13_bscale"][j].to(dev)
lo = lut[(pk & 0xF).long()]
hi = lut[(pk >> 4).long()]
w13 = torch.stack([lo, hi], -1).reshape(pk.shape[0], -1)
w13 *= bs.view(torch.float8_e4m3fn).float().repeat_interleave(16, -1)
w13[:2048] *= float(nvp["nvfp4_w13_scale2"][j, 0])
w13[2048:] *= float(nvp["nvfp4_w13_scale2"][j, 1])
pk2 = nvp["nvfp4_w2_packed"][j].to(dev)
bs2 = nvp["nvfp4_w2_bscale"][j].to(dev)
lo2 = lut[(pk2 & 0xF).long()]
hi2 = lut[(pk2 >> 4).long()]
w2 = torch.stack([lo2, hi2], -1).reshape(pk2.shape[0], -1)
w2 *= bs2.view(torch.float8_e4m3fn).float().repeat_interleave(16, -1)
w2 *= float(nvp["nvfp4_w2_scale2"][j, 0])
return w13, w2
def encode_scaled(w, h, cent):
"""2-sweep h-weighted encode with scale refit; cent frozen."""
M, K = w.shape
NG = K // G
scales = w.abs().mean(-1).clamp_min(1e-8)
centb = cent.t().to(torch.bfloat16)
cent2b = (cent.t() ** 2).to(torch.bfloat16)
hw = h.reshape(1, NG, G).expand(M, NG, G).reshape(-1, G)
codes = torch.empty(M * NG, dtype=torch.int32, device=dev)
for _ in range(2):
tgt = (w / scales.unsqueeze(-1)).reshape(-1, G)
hwb = hw.to(torch.bfloat16)
for s in range(0, tgt.shape[0], CHUNK):
v = tgt[s:s+CHUNK].to(torch.bfloat16)
ww = hwb[s:s+CHUNK]
a = ww @ cent2b
a -= 2 * ((v * ww) @ centb)
codes[s:s+CHUNK] = a.argmin(-1).to(torch.int32)
dec = cent[codes.long()].reshape(M, K)
num = (w * h * dec).sum(-1)
den = ((dec * dec) * h).sum(-1).clamp_min(1e-10)
scales = (num / den).clamp(1e-6, None)
c16 = torch.where(codes >= 32768, codes - 65536, codes).to(torch.int16)
return c16.reshape(1, M, NG), scales
cb13 = part["w13_codebooks"][0].float().to(dev)
cb2 = part["w2c_codebooks"][0].float().to(dev)
add = {"ids": [], "w13c": [], "w13s": [], "w2c": [], "w2s": []}
for e in new:
w13, w2 = teacher(e)
rows = (tk == e).any(1)
if int(rows.sum()) >= 16:
xe = x[rows]
h13 = xe.pow(2).mean(0).clamp_min(1e-10)
gate, up = w13[:2048], w13[2048:]
mid = torch.nn.functional.silu(xe @ gate.t()) * (xe @ up.t())
h2 = mid.pow(2).mean(0).clamp_min(1e-10)
else:
h13 = x.pow(2).mean(0).clamp_min(1e-10)
h2 = torch.ones(2048, device=dev)
c13, s13 = encode_scaled(w13, h13, cb13)
c2, s2 = encode_scaled(w2, h2, cb2)
add["ids"].append(e)
add["w13c"].append(c13.unsqueeze(0).cpu())
add["w13s"].append(s13.half().unsqueeze(0).cpu())
add["w2c"].append(c2.unsqueeze(0).cpu())
add["w2s"].append(s2.half().unsqueeze(0).cpu())
import torch as T
old_ids = part["expert_ids"].tolist()
all_ids = old_ids + add["ids"]
order = sorted(range(len(all_ids)), key=lambda i: all_ids[i])
cat = {
"w13_codes": T.cat([part["w13_codes"]] + add["w13c"]),
"w13_scales": T.cat([part["w13_scales"]] + add["w13s"]),
"w2c_codes": T.cat([part["w2c_codes"]] + add["w2c"]),
"w2c_scales": T.cat([part["w2c_scales"]] + add["w2s"]),
}
sel = T.tensor(order)
T.save({
"layer": li,
"expert_ids": T.tensor([all_ids[i] for i in order],
dtype=T.int32),
"w13_codes": cat["w13_codes"][sel].contiguous(),
"w13_codebooks": part["w13_codebooks"],
"w13_scales": cat["w13_scales"][sel].contiguous(),
"w2c_codes": cat["w2c_codes"][sel].contiguous(),
"w2c_codebooks": part["w2c_codebooks"],
"w2c_scales": cat["w2c_scales"][sel].contiguous(),
"w13_err_before": part["w13_err_before"],
"w13_err_after": part["w13_err_after"],
"w2_err_before": part["w2_err_before"],
"w2_err_after": part["w2_err_after"],
}, outp)
return li, f"+{len(new)} demoted encoded ({time.time()-t0:.0f}s)"
def main():
os.makedirs(OUT, exist_ok=True)
layers = range(3, 78)
if os.environ.get("P15_NEED"):
layers = sorted(int(k) for k in
json.load(open(os.environ["P15_NEED"])))
jobs = [(li, i % 8) for i, li in enumerate(layers)]
with ProcessPoolExecutor(max_workers=8) as ex:
for f in [ex.submit(process_layer, *j) for j in jobs]:
li, msg = f.result()
print(f"L{li}: {msg}", flush=True)
if __name__ == "__main__":
main()