McBopomofoLM-models / scripts /enc /extract_weights.py
workfunction's picture
McBopomofoLM v2.1.1 Core ML models (SlothE-T 25M encoder, pred_q35_60m decoder)
ff5f59d
Raw
History Blame Contribute Delete
3.61 kB
"""Extract SlothE-T 12M deploy weights from the fp32 master safetensors (upstream quantizer:
weight_quant_ternary median, per output channel) -> work/w12m.npz, and cross-check the ternary
codes/scales and the fp islands against the shipped TQ2_0/F16 GGUF that libslothe loads."""
import glob, json, os, sys, argparse
import numpy as np, torch
from safetensors.torch import load_file
import gguf
from gguf.quants import dequantize
ap = argparse.ArgumentParser(); ap.add_argument("--size", default="12m", choices=["12m", "25m"]); SZ = ap.parse_args().size
F = {"12m": ("12m/model.safetensors", "12m/slothe_config.json", "slothe-t-12m-256x12.gguf"),
"25m": ("model.safetensors", "slothe_config.json", "slothe-t-25m.gguf")}[SZ]
SNAP = glob.glob(os.path.expanduser("~/.cache/huggingface/hub/models--Luigi--sloth-ime-models/snapshots/e7d13c9c*/"))[0]
cfg = json.load(open(SNAP + F[1]))
sd = {k: v.float() for k, v in load_file(SNAP + F[0]).items()}
D = cfg["depth"]
out = {}
def tern(w): # verbatim math of slothe_torch.weight_quant_ternary(mode="median")
s = w.abs().median(dim=1, keepdim=True).values.clamp_(min=1e-5)
q = (w / s).round().clamp_(-1, 1)
return q.numpy().astype(np.int8), s[:, 0].numpy()
out["embed"] = sd["embed.weight"].numpy(); out["embed_norm"] = sd["embed_norm.w"].numpy()
out["norm"] = sd["norm.w"].numpy(); out["head"] = sd["head.weight"].numpy()
LIN = {"q": "attn.q", "k": "attn.k", "v": "attn.v", "o": "attn.o", "w1": "ffn.w1", "w3": "ffn.w3", "w2": "ffn.w2"}
for i in range(D):
fp = i == 0 or i == D - 1
p = f"blocks.{i}."
out[f"{i}.n1"] = sd[p + "n1.w"].numpy(); out[f"{i}.n2"] = sd[p + "n2.w"].numpy()
out[f"{i}.qn"] = sd[p + "attn.qn.w"].numpy(); out[f"{i}.kn"] = sd[p + "attn.kn.w"].numpy()
for s, n in LIN.items():
w = sd[p + n + ".weight"]
if fp:
out[f"{i}.{s}.w"] = w.numpy()
else:
q, sc = tern(w)
out[f"{i}.{s}.q"] = q; out[f"{i}.{s}.s"] = sc
out[f"{i}.{s}.pre"] = sd[p + n + ".pre.w"].numpy()
np.savez(f"work/w{SZ}.npz", **out)
json.dump(cfg, open(f"work/cfg{SZ}.json", "w"))
# ---- cross-check against the GGUF libslothe loads ----
r = gguf.GGUFReader(SNAP + F[2])
G = {t.name: t for t in r.tensors}
def g(name):
t = G[name]; shp = [int(x) for x in reversed(t.shape)]
return dequantize(t.data, t.tensor_type).reshape(shp).astype(np.float32)
GN = {"q": "attn_q", "k": "attn_k", "v": "attn_v", "o": "attn_output", "w1": "ffn_gate", "w3": "ffn_up", "w2": "ffn_down"}
worst_t, worst_f, code_mis = 0.0, 0.0, 0
for i in range(D):
fp = i == 0 or i == D - 1
for s, gn in GN.items():
gw = g(f"blk.{i}.{gn}.weight")
if fp:
worst_f = max(worst_f, float(np.abs(gw - out[f"{i}.{s}.w"].astype(np.float16).astype(np.float32)).max()))
else:
mine = out[f"{i}.{s}.q"] * out[f"{i}.{s}.s"][:, None].astype(np.float16).astype(np.float32)
pad = gw[:, mine.shape[1]:]; assert not pad.any(), "nonzero GGUF in-feature padding"
gw = gw[:, :mine.shape[1]] # GGUF pads ternary in-features to a multiple of 256 with zeros
worst_t = max(worst_t, float(np.abs(gw - mine).max()))
code_mis += int((np.sign(gw) != out[f"{i}.{s}.q"]).sum())
hd = g("output.weight"); worst_f = max(worst_f, float(np.abs(hd - out["head"].astype(np.float16).astype(np.float32)).max()))
print(json.dumps({"ternary_max_abs_vs_gguf": worst_t, "ternary_code_mismatch": code_mis,
"fp_island_max_abs_vs_gguf(f16)": worst_f,
"n_params": int(sum(v.size for v in out.values()))}))