File size: 3,609 Bytes
ff5f59d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
"""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()))}))