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