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