| |
| """Assemble the hybrid NVFP4+AQLM GLM-5.2 checkpoint at /data/glm52. |
| |
| Sources: |
| - /tmp/glm52-dl/nvfp4-full: lukealonso/GLM-5.2-NVFP4 (all non-expert |
| tensors verbatim; routed-expert tensors only for plan.nvfp4_layers) |
| - /data/glm52-aqlm-parts/layer_N.pt: AQLM codes/codebooks/scales for |
| the remaining expert layers |
| |
| Output: sharded safetensors (~4GB each) + index + config.json with the |
| nvfp4_aqlm_hybrid quantization_config + tokenizer/aux files. |
| """ |
|
|
| import json |
| import os |
| import re |
| import shutil |
| import time |
|
|
| import torch |
| from safetensors import safe_open |
| from safetensors.torch import save_file |
|
|
| ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) |
| SRC = "/tmp/glm52-dl/nvfp4-full" |
| PARTS = "/data/glm52-aqlm-parts" |
| DST = "/data/glm52" |
| SHARD_BYTES = 4 << 30 |
|
|
| plan = json.load(open(os.path.join(ROOT, "hybrid_plan.json"))) |
| NVFP4_LAYERS = set(plan["nvfp4_layers"]) |
| AQLM_LAYERS = sorted(plan["aqlm_mixed_layers"]) + sorted(plan["aqlm_cold_layers"]) |
|
|
|
|
| def keep(name: str) -> bool: |
| m = re.match(r"model\.layers\.(\d+)\.mlp\.experts\.\d+\.", name) |
| if m: |
| return int(m.group(1)) in NVFP4_LAYERS |
| return True |
|
|
|
|
| class ShardWriter: |
| def __init__(self, dst: str): |
| self.dst = dst |
| self.cur: dict[str, torch.Tensor] = {} |
| self.cur_bytes = 0 |
| self.n = 0 |
| self.weight_map: dict[str, str] = {} |
| self.total = 0 |
| self.files: list[str] = [] |
|
|
| def add(self, name: str, tensor: torch.Tensor): |
| nb = tensor.numel() * tensor.element_size() |
| if self.cur_bytes + nb > SHARD_BYTES and self.cur: |
| self.flush() |
| self.cur[name] = tensor |
| self.cur_bytes += nb |
| self.total += nb |
|
|
| def flush(self): |
| if not self.cur: |
| return |
| self.n += 1 |
| fname = f"model-{self.n:05d}.safetensors" |
| save_file(self.cur, os.path.join(self.dst, fname)) |
| for k in self.cur: |
| self.weight_map[k] = fname |
| print(f" wrote {fname} ({self.cur_bytes/1e9:.2f} GB)", flush=True) |
| self.files.append(fname) |
| self.cur = {} |
| self.cur_bytes = 0 |
|
|
| def finalize(self): |
| self.flush() |
| |
| total_n = self.n |
| wm = {} |
| for i, fname in enumerate(self.files, 1): |
| new = f"model-{i:05d}-of-{total_n:05d}.safetensors" |
| os.rename( |
| os.path.join(self.dst, fname), os.path.join(self.dst, new) |
| ) |
| for k, v in self.weight_map.items(): |
| if v == fname: |
| wm[k] = new |
| idx = { |
| "metadata": {"total_size": self.total}, |
| "weight_map": wm, |
| } |
| json.dump( |
| idx, |
| open(os.path.join(self.dst, "model.safetensors.index.json"), "w"), |
| indent=0, |
| ) |
| print(f"index: {len(wm)} tensors, {self.total/1e9:.1f} GB total") |
|
|
|
|
| def main(): |
| os.makedirs(DST, exist_ok=True) |
|
|
| idx = json.load(open(os.path.join(SRC, "model.safetensors.index.json"))) |
| weight_map = idx["weight_map"] |
| shards = sorted(set(weight_map.values())) |
|
|
| writer = ShardWriter(DST) |
|
|
| |
| for shard in shards: |
| with safe_open(os.path.join(SRC, shard), framework="pt") as f: |
| names = [n for n in f.keys() if keep(n)] |
| if not names: |
| continue |
| print(f"{shard}: keeping {len(names)} tensors", flush=True) |
| for n in names: |
| writer.add(n, f.get_tensor(n)) |
|
|
| |
| layer_books = {} |
| mse_report = {} |
| for li in AQLM_LAYERS: |
| part = os.path.join(PARTS, f"layer_{li}.pt") |
| |
| while ( |
| not os.path.exists(part) or time.time() - os.path.getmtime(part) < 60 |
| ): |
| print(f"waiting for quantizer: layer {li} ...", flush=True) |
| time.sleep(30) |
| d = torch.load(part, map_location="cpu", weights_only=True) |
| p = f"model.layers.{li}.mlp.experts" |
| writer.add(f"{p}.w13_codes", d["w13_codes"]) |
| writer.add(f"{p}.w13_codebooks", d["w13_codebooks"]) |
| writer.add(f"{p}.w13_scales", d["w13_scales"]) |
| writer.add(f"{p}.w2_codes", d["w2_codes"]) |
| writer.add(f"{p}.w2_codebooks", d["w2_codebooks"]) |
| writer.add(f"{p}.w2_scales", d["w2_scales"]) |
| layer_books[str(li)] = { |
| "w13": d["books"]["w13"], |
| "w2": d["books"]["w2"], |
| } |
| mse_report[li] = (d["w13_rel_mse"], d["w2_rel_mse"]) |
| print( |
| f"layer {li}: aqlm books={d['books']} " |
| f"rel_mse w13={d['w13_rel_mse']:.4f} w2={d['w2_rel_mse']:.4f}", |
| flush=True, |
| ) |
|
|
| writer.finalize() |
|
|
| |
| |
| |
| |
| cfg = json.load(open(os.path.join(SRC, "config.json"))) |
| nvfp4_qc = cfg["quantization_config"] |
| cfg["quantization_config"] = { |
| "quant_method": "nvfp4_aqlm_hybrid", |
| "nvfp4": nvfp4_qc, |
| "aqlm": {"entries": 65536, "group_size": 8}, |
| "aqlm_layer_books": layer_books, |
| } |
| json.dump(cfg, open(os.path.join(DST, "config.json"), "w"), indent=2) |
|
|
| |
| for f in os.listdir(SRC): |
| if ( |
| f.endswith(".json") |
| and f not in ("config.json", "model.safetensors.index.json") |
| or f.endswith((".txt", ".jinja", ".py", ".md")) |
| ): |
| shutil.copy2(os.path.join(SRC, f), os.path.join(DST, f)) |
| print(f"copied {f}") |
|
|
| json.dump( |
| {str(k): v for k, v in mse_report.items()}, |
| open(os.path.join(DST, "aqlm_mse_report.json"), "w"), |
| indent=2, |
| ) |
| print("DONE:", DST) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|