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