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