File size: 5,256 Bytes
fdc6474 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 | #!/usr/bin/env python3
"""Convert whole-NVFP4 layers of /data/glm52 to per-expert hybrid, in place.
For each layer in the assignment that still has per-expert NVFP4 tensors in
the checkpoint: build the v3 hybrid tensor set (hyb_kind, compacted NVFP4
hot arrays, sliced AQLM base groups from /data/glm52-aqlm-parts), write it
as a patch shard, update index + config, then physically strip the stale
per-expert tensors (vLLM's loader reads shard files, not the index).
"""
import json
import os
import re
import sys
import torch
from safetensors import safe_open
from safetensors.torch import save_file
DST = "/data/glm52"
PARTS = "/data/glm52-aqlm-parts"
ASSIGN = "/data/glm52-expert-assignment.json"
assignment = {int(k): v for k, v in json.load(open(ASSIGN)).items()}
idx_path = os.path.join(DST, "model.safetensors.index.json")
idx = json.load(open(idx_path))
wm = idx["weight_map"]
cfg = json.load(open(os.path.join(DST, "config.json")))
books_cfg = cfg["quantization_config"]["aqlm_layer_books"]
# layers that need conversion: in assignment but not yet in config books
todo = sorted(li for li in assignment if str(li) not in books_cfg)
print("layers to convert:", todo)
if not todo:
sys.exit(0)
class Reader:
def __init__(self):
self._open = {}
def get(self, name):
shard = wm[name]
if shard not in self._open:
self._open[shard] = safe_open(
os.path.join(DST, shard), framework="pt"
)
return self._open[shard].get_tensor(name)
reader = Reader()
n_exp, inter, hidden = 256, 2048, 6144
for li in todo:
hot = sorted(assignment[li]["hot"])
cold = sorted(assignment[li]["cold"])
hot_s, cold_s = set(hot), set(cold)
base = [e for e in range(n_exp) if e not in hot_s and e not in cold_s]
b_all = sorted(set(base) | cold_s)
kind = torch.ones(n_exp, dtype=torch.int8)
for e in hot:
kind[e] = 0
for e in cold:
kind[e] = 2
part = torch.load(
os.path.join(PARTS, f"layer_{li}.pt"), map_location="cpu",
weights_only=True,
)
p = f"model.layers.{li}.mlp.experts"
b_idx = torch.tensor(b_all, dtype=torch.long)
base_idx = torch.tensor(base, dtype=torch.long)
cold_idx = torch.tensor(cold, dtype=torch.long)
tensors = {
f"{p}.hyb_kind": kind,
f"{p}.w13_codes": part["w13_codes"][b_idx].contiguous(),
f"{p}.w13_codebooks": part["w13_codebooks"].clone(),
f"{p}.w13_scales": part["w13_scales"][b_idx].contiguous(),
f"{p}.w2m_codes": part["w2_codes"][base_idx].contiguous(),
f"{p}.w2m_codebooks": part["w2_codebooks"].clone(),
f"{p}.w2m_scales": part["w2_scales"][base_idx].contiguous(),
f"{p}.w2c_codes": part["w2_codes"][cold_idx, :1].clone(),
f"{p}.w2c_codebooks": part["w2_codebooks"][:1].clone(),
f"{p}.w2c_scales": part["w2_scales"][cold_idx].contiguous(),
}
na = len(hot)
w13_packed = torch.empty(na, 2 * inter, hidden // 2, dtype=torch.uint8)
w13_bscale = torch.empty(na, 2 * inter, hidden // 16, dtype=torch.uint8)
w13_scale2 = torch.empty(na, 2, dtype=torch.float32)
w2_packed = torch.empty(na, hidden, inter // 2, dtype=torch.uint8)
w2_bscale = torch.empty(na, hidden, inter // 16, dtype=torch.uint8)
w2_scale2 = torch.empty(na, 1, dtype=torch.float32)
for j, e in enumerate(hot):
ep = f"{p}.{e}"
w13_packed[j, :inter] = reader.get(f"{ep}.gate_proj.weight")
w13_packed[j, inter:] = reader.get(f"{ep}.up_proj.weight")
w2_packed[j] = reader.get(f"{ep}.down_proj.weight")
w13_bscale[j, :inter] = reader.get(
f"{ep}.gate_proj.weight_scale").view(torch.uint8)
w13_bscale[j, inter:] = reader.get(
f"{ep}.up_proj.weight_scale").view(torch.uint8)
w2_bscale[j] = reader.get(
f"{ep}.down_proj.weight_scale").view(torch.uint8)
w13_scale2[j, 0] = reader.get(f"{ep}.gate_proj.weight_scale_2").float()
w13_scale2[j, 1] = reader.get(f"{ep}.up_proj.weight_scale_2").float()
w2_scale2[j, 0] = reader.get(f"{ep}.down_proj.weight_scale_2").float()
tensors.update({
f"{p}.nvfp4_w13_packed": w13_packed,
f"{p}.nvfp4_w13_bscale": w13_bscale,
f"{p}.nvfp4_w13_scale2": w13_scale2,
f"{p}.nvfp4_w2_packed": w2_packed,
f"{p}.nvfp4_w2_bscale": w2_bscale,
f"{p}.nvfp4_w2_scale2": w2_scale2,
})
fname = f"model-hybrid-patch-layer{li}.safetensors"
save_file(tensors, os.path.join(DST, fname))
stale = [n for n in wm
if re.match(rf"model\.layers\.{li}\.mlp\.experts\.\d+\.", n)]
for n in stale:
del wm[n]
for n in tensors:
wm[n] = fname
books_cfg[str(li)] = {
"n_nvfp4": na, "n_base": len(base), "n_cold": len(cold),
}
nb = sum(t.numel() * t.element_size() for t in tensors.values())
print(f"layer {li}: hot={na} base={len(base)} cold={len(cold)} "
f"patch={nb/1e9:.2f} GB, dropped {len(stale)} per-expert tensors",
flush=True)
reader._open.clear()
json.dump(idx, open(idx_path, "w"), indent=0)
json.dump(cfg, open(os.path.join(DST, "config.json"), "w"), indent=2)
print("index/config updated; run strip_stale.py next")
|