maniac-11111's picture
Mach-2 Additive Medium
b0234bb
Raw History Blame Contribute Delete
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)