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