Download decode.py from SyzygyResearch/Mach-2-Additive-Medium: direct link, hf CLI and curl.
- Browser
- Download file 17.8 kB
-
https://huggingface.co/SyzygyResearch/Mach-2-Additive-Medium/resolve/main/decode.py
- Command line
-
hf download hf://SyzygyResearch/Mach-2-Additive-Medium/decode.py
-
curl -L -o decode.py https://huggingface.co/SyzygyResearch/Mach-2-Additive-Medium/resolve/main/decode.py
17.8 kB
| #!/usr/bin/env python3 | |
| """Decode the packed Mach-2 Additive Medium checkpoint into a bf16 Hugging Face checkpoint. | |
| python decode.py materialize --packed-dir . --out /path/ckpt [--workers 8] [--no-base] | |
| """ | |
| import argparse | |
| import hashlib | |
| import json | |
| import math | |
| import os | |
| import numpy as np | |
| HIDDEN, INTER, NEXP, NLAYERS = 2560, 640, 512, 48 | |
| BASE_REPO, BASE_REVISION = "Qwen/Qwen3.8-Flash-Next", "de4b8e4d43b917e7706784d8bb445c9af86a3540" | |
| FORMAT = "fn_packed_v2" | |
| def _hmat(q_elems, chi, sub): | |
| n = len(q_elems) | |
| Q = np.array([[chi(sub(q_elems[i], q_elems[j])) for j in range(n)] for i in range(n)], np.float64) | |
| C = np.zeros((n + 1, n + 1)) | |
| C[0, 1:] = 1.0 | |
| C[1:, 0] = 1.0 | |
| C[1:, 1:] = Q | |
| A = np.array([[1.0, 1.0], [1.0, -1.0]]) | |
| B = np.array([[1.0, -1.0], [-1.0, -1.0]]) | |
| H = np.kron(C, A) + np.kron(np.eye(n + 1), B) | |
| r = 2 * (n + 1) | |
| assert np.array_equal(H, H.T) and np.array_equal(H @ H, r * np.eye(r)) | |
| return H.astype(np.float32) | |
| def hadamard_12(): | |
| chi5 = {0: 0, 1: 1, 2: -1, 3: -1, 4: 1} | |
| return _hmat(list(range(5)), lambda e: chi5[e], lambda u, v: (u - v) % 5) | |
| def hadamard_20(): | |
| els = [(a, b) for a in range(3) for b in range(3)] | |
| def mul(u, v): | |
| (a, b), (c, d) = u, v | |
| return ((a * c - b * d) % 3, (a * d + b * c) % 3) | |
| squares = {mul(e, e) for e in els if e != (0, 0)} | |
| return _hmat(els, lambda e: 0 if e == (0, 0) else (1 if e in squares else -1), | |
| lambda u, v: ((u[0] - v[0]) % 3, (u[1] - v[1]) % 3)) | |
| _HS = { | |
| 12: ["+-----------", "++-+---+++-+", "+++-+---+++-", "+-++-+---+++", "++-++-+---++", "+++-++-+---+", | |
| "++++-++-+---", "+-+++-++-+--", "+--+++-++-+-", "+---+++-++-+", "++---+++-++-", "+-+---+++-++"], | |
| 20: ["+-------------------", "++-++----+-+-++++--+", "+++-++----+-+-++++--", "+-++-++----+-+-++++-", | |
| "+--++-++----+-+-++++", "++--++-++----+-+-+++", "+++--++-++----+-+-++", "++++--++-++----+-+-+", | |
| "+++++--++-++----+-+-", "+-++++--++-++----+-+", "++-++++--++-++----+-", "+-+-++++--++-++----+", | |
| "++-+-++++--++-++----", "+-+-+-++++--++-++---", "+--+-+-++++--++-++--", "+---+-+-++++--++-++-", | |
| "+----+-+-++++--++-++", "++----+-+-++++--++-+", "+++----+-+-++++--++-", "+-++----+-+-++++--++"], | |
| } | |
| _HR = {} | |
| def _hr(kind, radix): | |
| if (kind, radix) not in _HR: | |
| if kind == "spine": | |
| H = np.array([[1.0 if c == "+" else -1.0 for c in r] for r in _HS[radix]], np.float32) | |
| assert np.array_equal(H @ H.T, radix * np.eye(radix, dtype=np.float32)) | |
| else: | |
| H = hadamard_12() if radix == 12 else hadamard_20() | |
| _HR[(kind, radix)] = H | |
| return _HR[(kind, radix)] | |
| def hadamard(x, kind="expert"): | |
| x = np.asarray(x, np.float32) | |
| shape, N = x.shape, x.shape[-1] | |
| if N & (N - 1): | |
| radix = 12 if (N % 12 == 0 and (N // 12) & (N // 12 - 1) == 0) else 20 | |
| M = N // radix | |
| assert N % radix == 0 and M & (M - 1) == 0, f"Hadamard needs 2^k, 12*2^k or 20*2^k, got {N}" | |
| if kind == "spine": | |
| xb = _butterflies(x.reshape(-1, radix, M)) if M > 1 else x.reshape(-1, radix, 1) | |
| return (np.matmul(_hr(kind, radix), xb) / np.float32(math.sqrt(N))).reshape(shape) | |
| xb = hadamard(x.reshape(-1, radix, M)) if M > 1 else x.reshape(-1, radix, 1) | |
| return (np.matmul(_hr(kind, radix), xb) / np.float32(math.sqrt(radix))).reshape(shape) | |
| return (_butterflies(x) / np.float32(math.sqrt(N))).reshape(shape) | |
| def _butterflies(x): | |
| N = x.shape[-1] | |
| cur = np.ascontiguousarray(x, np.float32).reshape(-1, N) | |
| span = 1 | |
| while span < N: | |
| blk = cur.reshape(-1, N // (2 * span), 2, span) | |
| cur = np.stack([blk[:, :, 0] + blk[:, :, 1], blk[:, :, 0] - blk[:, :, 1]], 2).reshape(-1, N) | |
| span *= 2 | |
| return cur.reshape(x.shape) | |
| def bf16_round(x): | |
| u = np.ascontiguousarray(x, np.float32).view(np.uint32) | |
| return ((u + np.uint32(0x7FFF) + ((u >> np.uint32(16)) & np.uint32(1))) & np.uint32(0xFFFF0000)).view(np.float32) | |
| def _read(path): | |
| from safetensors import safe_open | |
| with safe_open(path, framework="np") as fh: | |
| return {k: fh.get_tensor(k) for k in fh.keys()}, dict(fh.metadata() or {}) | |
| EX_V, EX_L, EX_TD = 4, 16, 16 | |
| EX_SHAPES = {"gate": (INTER, HIDDEN), "up": (INTER, HIDDEN), "down": (HIDDEN, INTER)} | |
| def _ex_states(words, K4): | |
| words = np.asarray(words).view(np.uint16).astype(np.int64) | |
| rows, nstep, step = words.shape[0], EX_TD * EX_TD // EX_V, int(K4) | |
| bits = ((words[:, :, None] >> np.arange(15, -1, -1)) & 1).reshape(rows, -1)[:, :step * nstep] | |
| bits = np.concatenate([bits, bits[:, :EX_L - step]], axis=1) | |
| w = 1 << np.arange(EX_L - 1, -1, -1, dtype=np.int64) | |
| idx = np.arange(nstep)[:, None] * step + np.arange(EX_L)[None, :] | |
| return bits[:, idx] @ w | |
| def wave_map(Mb, Nb): | |
| starts = [(Mb - i - 1, Nb - 1) for i in range(Mb)] + [(0, Nb - i - 1) for i in range(Nb)] | |
| idx = np.zeros((Mb, Nb), dtype=np.int64) | |
| for w, (jm, jn) in enumerate(starts): | |
| while 0 <= jm < Mb and 0 <= jn < Nb: | |
| idx[jm, jn] = w | |
| jm, jn = jm + 1, jn - 1 | |
| return idx | |
| def sign_seed(layer, expert, proj, side, ns="canon"): | |
| return int.from_bytes(hashlib.sha256(f"{ns}|L{layer}|e{expert}|{proj}|{side}".encode()).digest()[:7], "big") | |
| _PM0, _PM1 = np.uint64(0xD2511F53), np.uint64(0xCD9E8D57) | |
| _PW0, _PW1, _PMASK = np.uint64(0x9E3779B9), np.uint64(0xBB67AE85), np.uint64(0xFFFFFFFF) | |
| def expert_signs(seeds, dim): | |
| E = len(seeds) | |
| c2 = np.tile(np.arange(dim, dtype=np.uint64), E) | |
| k0 = np.repeat(np.array([s & 0xFFFFFFFF for s in seeds], np.uint64), dim) | |
| k1 = np.repeat(np.array([s >> 32 for s in seeds], np.uint64), dim) | |
| c0 = c1 = c3 = np.zeros(E * dim, np.uint64) | |
| for r in range(10): | |
| if r: | |
| k0, k1 = (k0 + _PW0) & _PMASK, (k1 + _PW1) & _PMASK | |
| p0, p1 = _PM0 * c0, _PM1 * c2 | |
| c0, c1, c2, c3 = (p1 >> np.uint64(32)) ^ c1 ^ k0, p1 & _PMASK, (p0 >> np.uint64(32)) ^ c3 ^ k1, p0 & _PMASK | |
| inv2pi = np.float32(2 * np.pi / 2 ** 32) | |
| v = (c1.astype(np.float32) * inv2pi + inv2pi / np.float32(2)).astype(np.float32) | |
| u = (c0.astype(np.float32) * np.float32(2.0 ** -32) + np.float32(2.0 ** -33)).astype(np.float32) | |
| return np.where((v <= np.float32(np.pi)) | (u == np.float32(1.0)), 1.0, -1.0).astype(np.float32).reshape(E, dim) | |
| def _lut(cb, K4): | |
| return np.asarray(cb.get(f"lut.k{K4}", cb["lut"]), np.float32) | |
| def decode_expert(t, cb, proj, e): | |
| m, n = EX_SHAPES[proj] | |
| K4 = int(t[f"{proj}.rate_k4"][e]) | |
| ex = np.asarray(t[f"k{K4}.{proj}.experts"]) | |
| row = int(np.searchsorted(ex, e)) | |
| assert row < ex.size and ex[row] == e, (proj, e, K4) | |
| states = _ex_states(t[f"k{K4}.{proj}.trellis"][row], K4) | |
| Mb, Nb = m // EX_TD, n // EX_TD | |
| unit = _lut(cb, K4)[states].reshape(Mb, Nb, EX_TD, EX_TD).transpose(0, 2, 1, 3).reshape(m, n) | |
| g = np.asarray(t[f"{proj}.wave_gamma"][e], np.float32)[wave_map(Mb, Nb)] | |
| unit = (unit.reshape(Mb, EX_TD, Nb, EX_TD) * g[:, None, :, None]).reshape(m, n) | |
| unit = unit * np.float32(t[f"{proj}.wscale"][e]) | |
| rows = hadamard(unit) * t[f"{proj}.SU"][e] | |
| cols = hadamard(np.ascontiguousarray(rows.T)) * t[f"{proj}.SV"][e] | |
| return np.ascontiguousarray(cols.T) | |
| def load_expert_layer(packed_dir, layer): | |
| t, meta = _read(os.path.join(packed_dir, "experts", f"L{layer:02d}.safetensors")) | |
| assert meta.get("format") == FORMAT, f"L{layer}: not an {FORMAT} expert shard ({meta.get('format')})" | |
| cb = _read(os.path.join(packed_dir, "experts", "codebook.safetensors"))[0] | |
| return with_signs(t, layer), cb | |
| def with_signs(t, layer): | |
| for p, (m, n) in EX_SHAPES.items(): | |
| if f"{p}.SU" not in t: | |
| t[f"{p}.SU"] = expert_signs([sign_seed(layer, e, p, "SU") for e in range(NEXP)], n) | |
| t[f"{p}.SV"] = expert_signs([sign_seed(layer, e, p, "SV") for e in range(NEXP)], m) | |
| return t | |
| def decode_expert_layer(packed_dir, layer, experts=None): | |
| t, cb = load_expert_layer(packed_dir, layer) | |
| experts = list(range(NEXP)) if experts is None else list(experts) | |
| gu = np.empty((len(experts), 2 * INTER, HIDDEN), np.float32) | |
| dn = np.empty((len(experts), HIDDEN, INTER), np.float32) | |
| for i, e in enumerate(experts): | |
| gu[i, :INTER] = decode_expert(t, cb, "gate", e) | |
| gu[i, INTER:] = decode_expert(t, cb, "up", e) | |
| dn[i] = decode_expert(t, cb, "down", e) | |
| return {"gate_up_proj": gu, "down_proj": dn} | |
| NE_K, NE_L, NE_V, NE_TLUT_BITS, NE_TD = 4, 16, 2, 9, 16 | |
| _FULL_LUT = {} | |
| def _ne_full_lut(tlut): | |
| key = np.asarray(tlut).tobytes() | |
| if key not in _FULL_LUT: | |
| small = np.asarray(tlut, np.float32) | |
| s = np.arange(1 << NE_L, dtype=np.int64) | |
| p = s * (s + 1) | |
| row = (p >> (16 - NE_TLUT_BITS - 1)) & ((1 << NE_TLUT_BITS) - 1) | |
| table = small[row].copy() | |
| table[:, 0] *= (1 - ((p >> 15) & 1) * 2).astype(np.float32) | |
| _FULL_LUT[key] = table | |
| return _FULL_LUT[key] | |
| def _ne_states(words): | |
| words = np.asarray(words).view(np.uint16).astype(np.int64) | |
| rows, T = words.shape[0], NE_TD * NE_TD | |
| step, nstep = NE_K * NE_V, T // NE_V | |
| bits = ((words[:, :, None] >> np.arange(15, -1, -1)) & 1).reshape(rows, -1)[:, :T * NE_K] | |
| bits = np.concatenate([bits, bits[:, :NE_L - step]], axis=1) | |
| w = 1 << np.arange(NE_L - 1, -1, -1, dtype=np.int64) | |
| idx = np.arange(nstep)[:, None] * step + np.arange(NE_L)[None, :] | |
| return bits[:, idx] @ w | |
| def decode_ne_tensor(t, name, m, n, tlut): | |
| states = _ne_states(t[f"{name}|trellis"]) | |
| vals = _ne_full_lut(tlut)[states] | |
| unit = np.ascontiguousarray(vals.reshape(m // NE_TD, n // NE_TD, NE_TD, NE_TD).transpose(0, 2, 1, 3)).reshape(m, n) | |
| unit = unit * np.float32(np.asarray(t[f"{name}|Wscale"]).reshape(-1)[0]) | |
| su, sv = np.asarray(t[f"{name}|SU"]), np.asarray(t[f"{name}|SV"]) | |
| rows = hadamard(unit, "spine") * np.sign(su).astype(np.float32) | |
| cols = hadamard(np.ascontiguousarray(rows.T), "spine") * np.sign(sv).astype(np.float32) | |
| w = np.ascontiguousarray(cols.T) | |
| mx = np.asarray(t[f"{name}|rc_max"], np.float32) | |
| r, c = rc_grid(mx[0], np.abs(sv.astype(np.int32))), rc_grid(mx[1], np.abs(su.astype(np.int32))) | |
| return (r[:, None] * bf16_round(w)) * c[None, :] | |
| def rc_grid(mx, k): | |
| return (np.float32(mx) * k.astype(np.float32)).astype(np.float32) * np.float32(1.0 / 127.0) | |
| def decode_ne_shard(packed_dir, layer): | |
| t, meta = _read(os.path.join(packed_dir, "ne", f"L{layer:02d}.safetensors")) | |
| tlut = _read(os.path.join(packed_dir, "ne", "tlut.safetensors"))[0]["tlut"] | |
| dims = json.loads(meta["dims"]) | |
| return {name: decode_ne_tensor(t, name, d[2], d[3], tlut)[:d[0], :d[1]] for name, d in dims.items()} | |
| def unpack_int5(qp, n): | |
| m = qp.shape[0] | |
| by = np.asarray(qp).reshape(m, n // 8, 5) | |
| full = np.zeros((m, n // 8, 8), dtype=np.uint8) | |
| full[:, :, :5] = by | |
| word = full.reshape(m, n).view("<u8").reshape(m, n // 8) | |
| out = np.zeros((m, n // 8, 8), dtype=np.int8) | |
| for i in range(8): | |
| out[:, :, i] = ((word >> np.uint64(5 * i)) & np.uint64(31)).astype(np.int8) - 16 | |
| return out.reshape(m, n) | |
| def decode_head(packed_dir): | |
| d = os.path.join(packed_dir, "head") | |
| parts = [] | |
| for f in sorted(x for x in os.listdir(d) if x.startswith("head_c") and x.endswith(".safetensors")): | |
| t, meta = _read(os.path.join(d, f)) | |
| parts += [(int(name.split(":")[1]), m0, n0, int(meta.get("group", 64)), t, name) | |
| for name, (m0, n0) in json.loads(meta["dims"]).items()] | |
| parts.sort(key=lambda x: x[0]) | |
| out = np.empty((sum(p[1] for p in parts), parts[0][2]), np.float32) | |
| for r0, m0, n0, g, t, name in parts: | |
| q = unpack_int5(t[f"{name}|qp"], n0).astype(np.float32) | |
| out[r0:r0 + m0] = q * np.repeat(np.asarray(t[f"{name}|gscale"], np.float32), g, axis=1)[:, :n0] | |
| return out | |
| def decode_embed(packed_dir, bits=4, rows_per=16384): | |
| t, _ = _read(os.path.join(packed_dir, "ne", f"embed_int{bits}.safetensors")) | |
| rows, ng = t["mn"].shape | |
| hid = t["q_packed"].shape[1] * 8 // bits | |
| out = np.empty((rows, hid), np.float32) | |
| for r0 in range(0, rows, rows_per): | |
| sl = slice(r0, min(r0 + rows_per, rows)) | |
| b = np.unpackbits(t["q_packed"][sl], axis=1, count=hid * bits).reshape(-1, hid, bits) | |
| q = np.zeros(b.shape[:2], np.uint8) | |
| for j in range(bits): | |
| q = (q << 1) | b[..., j] | |
| mn = t["mn"][sl].astype(np.float32)[..., None] | |
| mx = t["mx"][sl].astype(np.float32)[..., None] | |
| step = np.maximum(mx - mn, np.float32(1e-8)) * np.float32(1.0 / (2 ** bits - 1)) | |
| out[sl] = (mn + q.reshape(-1, ng, hid // ng).astype(np.float32) * step).reshape(-1, hid) | |
| return out | |
| def decode_int8_rows(packed_dir): | |
| from safetensors import safe_open | |
| out = {} | |
| with safe_open(os.path.join(packed_dir, "ne", "int8_rows.safetensors"), framework="pt") as fh: | |
| names = sorted({k.split("|")[0] for k in fh.keys()}) | |
| for name in names: | |
| q = fh.get_tensor(f"{name}|q").numpy().astype(np.float32) | |
| s = fh.get_tensor(f"{name}|scale").float().numpy() | |
| out[name] = q * s | |
| return out | |
| def _bf16(a): | |
| import torch | |
| return torch.from_numpy(np.ascontiguousarray(a, np.float32)).to(torch.bfloat16) | |
| def _expert_layer_job(args): | |
| packed_dir, out, L = args | |
| import torch | |
| from safetensors.torch import save_file | |
| t, cb = load_expert_layer(packed_dir, L) | |
| gu = torch.empty((NEXP, 2 * INTER, HIDDEN), dtype=torch.bfloat16) | |
| dn = torch.empty((NEXP, HIDDEN, INTER), dtype=torch.bfloat16) | |
| for e in range(NEXP): | |
| gu[e, :INTER] = _bf16(decode_expert(t, cb, "gate", e)) | |
| gu[e, INTER:] = _bf16(decode_expert(t, cb, "up", e)) | |
| dn[e] = _bf16(decode_expert(t, cb, "down", e)) | |
| p = f"model.language_model.layers.{L}.mlp.experts." | |
| fn = f"experts-L{L:02d}.safetensors" | |
| save_file({p + "gate_up_proj": gu, p + "down_proj": dn}, os.path.join(out, fn), metadata={"format": "pt"}) | |
| return {p + "gate_up_proj": fn, p + "down_proj": fn} | |
| def _ne_layer_job(args): | |
| packed_dir, out, L = args | |
| from safetensors.torch import save_file | |
| dec = decode_ne_shard(packed_dir, L) | |
| fn = f"spine-L{L:02d}.safetensors" | |
| save_file({k: _bf16(v) for k, v in dec.items()}, os.path.join(out, fn), metadata={"format": "pt"}) | |
| return {k: fn for k in dec} | |
| def materialize(packed_dir, out, workers=4, base=True): | |
| import shutil | |
| from concurrent.futures import ProcessPoolExecutor | |
| from safetensors import safe_open | |
| from safetensors.torch import save_file | |
| os.makedirs(out, exist_ok=True) | |
| root = os.path.dirname(os.path.abspath(packed_dir.rstrip("/"))) if os.path.basename(packed_dir.rstrip("/")) == "packed" \ | |
| else packed_dir | |
| pk = os.path.join(root, "packed") | |
| wm = {} | |
| with ProcessPoolExecutor(workers) as ex: | |
| for r in ex.map(_ne_layer_job, [(pk, out, L) for L in range(NLAYERS)]): | |
| wm.update(r) | |
| for r in ex.map(_expert_layer_job, [(pk, out, L) for L in range(NLAYERS)]): | |
| wm.update(r) | |
| save_file({"lm_head.weight": _bf16(decode_head(pk))}, os.path.join(out, "head.safetensors"), metadata={"format": "pt"}) | |
| wm["lm_head.weight"] = "head.safetensors" | |
| emb = "model.language_model.embed_tokens.weight" | |
| save_file({emb: _bf16(decode_embed(pk))}, os.path.join(out, "embed.safetensors"), metadata={"format": "pt"}) | |
| wm[emb] = "embed.safetensors" | |
| i8 = decode_int8_rows(pk) | |
| save_file({k: _bf16(v) for k, v in i8.items()}, os.path.join(out, "int8rows.safetensors"), metadata={"format": "pt"}) | |
| wm.update({k: "int8rows.safetensors" for k in i8}) | |
| for f in ["extras.safetensors"] + sorted(os.path.join("packed", "table", x) for x in os.listdir(os.path.join(pk, "table")) | |
| if x.endswith(".safetensors")): | |
| dst = os.path.basename(f) | |
| shutil.copyfile(os.path.join(root, f), os.path.join(out, dst)) | |
| with safe_open(os.path.join(out, dst), framework="pt") as fh: | |
| wm.update({k: dst for k in fh.keys()}) | |
| if base: | |
| from huggingface_hub import hf_hub_download | |
| idx = json.load(open(hf_hub_download(BASE_REPO, "model.safetensors.index.json", revision=BASE_REVISION)))["weight_map"] | |
| want = {k: f for k, f in idx.items() if k.startswith("mtp.") or k.startswith("model.visual.")} | |
| for f in sorted(set(want.values())): | |
| src = hf_hub_download(BASE_REPO, f, revision=BASE_REVISION) | |
| with safe_open(src, framework="pt") as fh: | |
| ks = [k for k in fh.keys() if k in want] | |
| save_file({k: fh.get_tensor(k) for k in ks}, os.path.join(out, f"base-{f}"), metadata={"format": "pt"}) | |
| wm.update({k: f"base-{f}" for k in ks}) | |
| for f in os.listdir(root): | |
| if f.endswith((".json", ".jinja", ".txt")) and f not in ("MANIFEST.json", "model.safetensors.index.json"): | |
| shutil.copyfile(os.path.join(root, f), os.path.join(out, f)) | |
| json.dump({"metadata": {}, "weight_map": dict(sorted(wm.items()))}, open(os.path.join(out, "model.safetensors.index.json"), "w"), | |
| indent=1) | |
| print(f"MATERIALIZED {len(wm)} tensors -> {out}", flush=True) | |
| if __name__ == "__main__": | |
| ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0]) | |
| sub = ap.add_subparsers(dest="cmd", required=True) | |
| m = sub.add_parser("materialize") | |
| m.add_argument("--packed-dir", default=".") | |
| m.add_argument("--out", required=True) | |
| m.add_argument("--workers", type=int, default=4) | |
| m.add_argument("--no-base", action="store_true", help="text-only: skip MTP and vision weights") | |
| a = ap.parse_args() | |
| materialize(a.packed_dir, a.out, a.workers, not a.no_base) | |