| |
| """Re-tier a two-tier checkpoint per a REAP assignment (phase 1.5). |
| |
| Per hybrid layer of TARGET: |
| hyb_kind <- REAP assignment (0=hot NVFP4, 2=cold AQLM) |
| nvfp4_* <- retained-hot experts sliced from TARGET's own arrays; |
| promoted experts from the NVFP4 teacher regions |
| (/tmp/glm52-hot-dl2 + /tmp/glm52-hot-dl fallback) |
| w13_*/w2c_* <- sliced from /data/glm52-aqlm-conv15 (covers the union |
| of all budgets' cold sets) |
| w2m_* <- kept (empty) |
| Non-expert tensors verbatim. Writes TARGET-reap; caller gates + swaps. |
| |
| Usage: build_checkpoint_v8.py TARGET ASSIGNMENT_JSON |
| """ |
| import json |
| import os |
| import re |
| import shutil |
| import sys |
|
|
| import torch |
| from safetensors import safe_open |
| from safetensors.torch import save_file |
|
|
| TARGET = sys.argv[1].rstrip("/") |
| ASSIGN = sys.argv[2] |
| PARTS = "/data/glm52-aqlm-conv15" |
| DLS = ["/tmp/glm52-hot-dl2", "/tmp/glm52-hot-dl"] |
| DST = TARGET + "-reap" |
| SHARD_BYTES = 4 << 30 |
| N_EXP, INTER, HIDDEN = 256, 2048, 6144 |
|
|
| assignment = {int(k): v for k, v in json.load(open(ASSIGN)).items()} |
|
|
| DT = {"U8": torch.uint8, "F8_E4M3": torch.uint8, "BF16": torch.bfloat16, |
| "F32": torch.float32} |
|
|
|
|
| class Regions: |
| def __init__(self): |
| self.sources = [] |
| for dl in DLS: |
| if not os.path.exists(f"{dl}/headers.json"): |
| continue |
| headers = json.load(open(f"{dl}/headers.json")) |
| wm = {n: s for s, h in headers.items() |
| for n in h["header"] if n != "__metadata__"} |
| regions = {} |
| for shard in 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))) |
| regions[shard] = sorted(regs) |
| self.sources.append((headers, wm, regions)) |
|
|
| def get(self, name): |
| |
| |
| for headers, wm, regions in self.sources: |
| if name not in wm: |
| continue |
| shard = wm[name] |
| info = headers[shard]["header"][name] |
| b, e = info["data_offsets"] |
| for rb, path, sz in 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"]) |
| |
| if not hasattr(self, "_old_wm"): |
| oidx = json.load(open( |
| "/data/glm52-old-layerwise/model.safetensors.index.json")) |
| self._old_wm = oidx["weight_map"] |
| self._old_open = {} |
| if name in self._old_wm: |
| sh = self._old_wm[name] |
| if sh not in self._old_open: |
| self._old_open[sh] = safe_open( |
| f"/data/glm52-old-layerwise/{sh}", framework="pt") |
| return self._old_open[sh].get_tensor(name) |
| raise KeyError(name) |
|
|
|
|
| idx = json.load(open(f"{TARGET}/model.safetensors.index.json")) |
| wm = idx["weight_map"] |
| opened = {} |
|
|
|
|
| def get(name): |
| s = wm[name] |
| if s not in opened: |
| opened[s] = safe_open(f"{TARGET}/{s}", framework="pt") |
| return opened[s].get_tensor(name) |
|
|
|
|
| class Writer: |
| def __init__(self): |
| os.makedirs(DST, exist_ok=True) |
| self.cur, self.cur_bytes, self.n, self.total = {}, 0, 0, 0 |
| self.weight_map, self.files = {}, [] |
|
|
| def add(self, name, t): |
| nb = t.numel() * t.element_size() |
| if self.cur_bytes + nb > SHARD_BYTES and self.cur: |
| self.flush() |
| self.cur[name] = t |
| self.cur_bytes += nb |
| self.total += nb |
|
|
| def flush(self): |
| if not self.cur: |
| return |
| self.n += 1 |
| f = f"model-{self.n:05d}.safetensors" |
| save_file(self.cur, f"{DST}/{f}") |
| for k in self.cur: |
| self.weight_map[k] = f |
| self.files.append(f) |
| self.cur, self.cur_bytes = {}, 0 |
|
|
| def finalize(self): |
| self.flush() |
| out = {} |
| for i, f in enumerate(self.files, 1): |
| new = f"model-{i:05d}-of-{self.n:05d}.safetensors" |
| os.rename(f"{DST}/{f}", f"{DST}/{new}") |
| for k, v in self.weight_map.items(): |
| if v == f: |
| out[k] = new |
| json.dump({"metadata": {"total_size": self.total}, |
| "weight_map": out}, |
| open(f"{DST}/model.safetensors.index.json", "w"), indent=0) |
| print(f"index: {len(out)} tensors, {self.total/1e9:.1f} GB") |
|
|
|
|
| def layer_tensors(li, rr): |
| p = f"model.layers.{li}.mlp.experts" |
| hot = sorted(assignment[li]["hot"]) |
| cold = sorted(assignment[li]["cold"]) |
| assert len(hot) + len(cold) == N_EXP |
|
|
| kind_old = get(f"{p}.hyb_kind") |
| old_hot = (kind_old == 0).nonzero().flatten().tolist() |
| pos = {e: j for j, e in enumerate(old_hot)} |
| olds = {n: 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")} |
|
|
| na = len(hot) |
| w13p = torch.empty(na, 2*INTER, HIDDEN//2, dtype=torch.uint8) |
| w13b = torch.empty(na, 2*INTER, HIDDEN//16, dtype=torch.uint8) |
| w13s = torch.empty(na, 2, dtype=torch.float32) |
| w2p = torch.empty(na, HIDDEN, INTER//2, dtype=torch.uint8) |
| w2b = torch.empty(na, HIDDEN, INTER//16, dtype=torch.uint8) |
| w2s = torch.empty(na, 1, dtype=torch.float32) |
| n_promoted = 0 |
| for j, e in enumerate(hot): |
| if e in pos: |
| k = pos[e] |
| w13p[j] = olds["nvfp4_w13_packed"][k] |
| w13b[j] = olds["nvfp4_w13_bscale"][k] |
| w13s[j] = olds["nvfp4_w13_scale2"][k] |
| w2p[j] = olds["nvfp4_w2_packed"][k] |
| w2b[j] = olds["nvfp4_w2_bscale"][k] |
| w2s[j] = olds["nvfp4_w2_scale2"][k] |
| else: |
| n_promoted += 1 |
| ep = f"model.layers.{li}.mlp.experts.{e}" |
| w13p[j, :INTER] = rr.get(f"{ep}.gate_proj.weight") |
| w13p[j, INTER:] = rr.get(f"{ep}.up_proj.weight") |
| w2p[j] = rr.get(f"{ep}.down_proj.weight") |
| w13b[j, :INTER] = rr.get( |
| f"{ep}.gate_proj.weight_scale").view(torch.uint8) |
| w13b[j, INTER:] = rr.get( |
| f"{ep}.up_proj.weight_scale").view(torch.uint8) |
| w2b[j] = rr.get(f"{ep}.down_proj.weight_scale").view(torch.uint8) |
| w13s[j, 0] = rr.get(f"{ep}.gate_proj.weight_scale_2").float() |
| w13s[j, 1] = rr.get(f"{ep}.up_proj.weight_scale_2").float() |
| w2s[j, 0] = rr.get(f"{ep}.down_proj.weight_scale_2").float() |
|
|
| part = torch.load(f"{PARTS}/layer_{li}.pt", map_location="cpu", |
| weights_only=True) |
| ppos = {int(e): j for j, e in enumerate(part["expert_ids"].tolist())} |
| missing = [e for e in cold if e not in ppos] |
| assert not missing, f"L{li}: conv15 missing {missing[:5]}" |
| sel = torch.tensor([ppos[e] for e in cold], dtype=torch.long) |
|
|
| kind = torch.full((N_EXP,), 2, dtype=torch.int8) |
| for e in hot: |
| kind[e] = 0 |
| print(f"L{li}: hot={na} ({n_promoted} promoted) cold={len(cold)}", |
| flush=True) |
| return { |
| "hyb_kind": kind, |
| "w13_codes": part["w13_codes"][sel].contiguous(), |
| "w13_codebooks": part["w13_codebooks"].clone(), |
| "w13_scales": part["w13_scales"][sel].contiguous(), |
| "w2c_codes": part["w2c_codes"][sel].contiguous(), |
| "w2c_codebooks": part["w2c_codebooks"].clone(), |
| "w2c_scales": part["w2c_scales"][sel].contiguous(), |
| "nvfp4_w13_packed": w13p, "nvfp4_w13_bscale": w13b, |
| "nvfp4_w13_scale2": w13s, "nvfp4_w2_packed": w2p, |
| "nvfp4_w2_bscale": w2b, "nvfp4_w2_scale2": w2s, |
| }, {"n_nvfp4": na, "n_base": 0, "n_cold": len(cold)} |
|
|
|
|
| REPL = ("hyb_kind", "w13_codes", "w13_codebooks", "w13_scales", |
| "w2c_codes", "w2c_codebooks", "w2c_scales", "nvfp4_w13_packed", |
| "nvfp4_w13_bscale", "nvfp4_w13_scale2", "nvfp4_w2_packed", |
| "nvfp4_w2_bscale", "nvfp4_w2_scale2") |
| pat = re.compile(r"model\.layers\.(\d+)\.mlp\.experts\.(" |
| + "|".join(REPL) + r")$") |
|
|
| rr = Regions() |
| w = Writer() |
| cache_li, cache = None, None |
| books = {} |
| for shard in sorted(set(wm.values())): |
| with safe_open(f"{TARGET}/{shard}", framework="pt") as f: |
| for name in f.keys(): |
| m = pat.match(name) |
| if m and int(m.group(1)) in assignment: |
| li = int(m.group(1)) |
| if cache_li != li: |
| cache, b = layer_tensors(li, rr) |
| books[str(li)] = b |
| cache_li = li |
| w.add(name, cache[m.group(2)]) |
| else: |
| w.add(name, f.get_tensor(name)) |
| w.finalize() |
|
|
| cfg = json.load(open(f"{TARGET}/config.json")) |
| cfg["quantization_config"]["aqlm_layer_books"] = books |
| json.dump(cfg, open(f"{DST}/config.json", "w"), indent=2) |
| for f in os.listdir(TARGET): |
| if (f.endswith(".json") and f not in ("config.json", "model.safetensors.index.json") |
| or f.endswith((".txt", ".jinja", ".py", ".md", ".sh"))): |
| shutil.copy2(f"{TARGET}/{f}", f"{DST}/{f}") |
| print("DONE:", DST) |
|
|