#!/usr/bin/env python3 """Replace a checkpoint's cold-expert AQLM arrays with converged parts. Streams every shard of TARGET; copies all tensors verbatim EXCEPT model.layers.N.mlp.experts.{w13_codes,w13_codebooks,w13_scales, w2c_codes,w2c_codebooks,w2c_scales} for hybrid layers, which are replaced by slicing /data/glm52-aqlm-conv/layer_N.pt (fit over the 1M cold-set superset) down to the TARGET's own cold ids (from its hyb_kind). w2m_* stay (empty in two-tier checkpoints). Writes TARGET-conv; caller gates it and swaps. Usage: build_checkpoint_v7.py TARGET [PARTS_DIR] """ 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("/") PARTS = sys.argv[2] if len(sys.argv) > 2 else "/data/glm52-aqlm-conv" DST = TARGET + "-conv" SHARD_BYTES = 4 << 30 REPLACE = ("w13_codes", "w13_codebooks", "w13_scales", "w2c_codes", "w2c_codebooks", "w2c_scales") pat = re.compile( r"model\.layers\.(\d+)\.mlp\.experts\.(" + "|".join(REPLACE) + r")$") 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 sliced(li): """Per-layer replacement tensors sliced to the target's cold ids.""" kind = get(f"model.layers.{li}.mlp.experts.hyb_kind") cold = (kind == 2).nonzero().flatten().tolist() part = torch.load(f"{PARTS}/layer_{li}.pt", map_location="cpu", weights_only=True) pos = {int(e): j for j, e in enumerate(part["expert_ids"].tolist())} missing = [e for e in cold if e not in pos] assert not missing, f"L{li}: parts missing cold ids {missing[:5]}" sel = torch.tensor([pos[e] for e in cold], dtype=torch.long) return { "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(), } w = Writer() cache = {} replaced = 0 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: li = int(m.group(1)) if li not in cache: cache = {li: sliced(li)} # one layer resident at a time w.add(name, cache[li][m.group(2)]) replaced += 1 else: w.add(name, f.get_tensor(name)) w.finalize() print(f"replaced {replaced} tensors from {PARTS}") for f in os.listdir(TARGET): if (f.endswith(".json") and f != "model.safetensors.index.json" or f.endswith((".txt", ".jinja", ".py", ".md", ".sh"))): shutil.copy2(f"{TARGET}/{f}", f"{DST}/{f}") print("DONE:", DST)