workfunction's picture
McBopomofoLM v2.1.1 Core ML models (SlothE-T 25M encoder, pred_q35_60m decoder)
ff5f59d
Raw
History Blame Contribute Delete
2.78 kB
"""numpy fp32 forward of SlothE-T 12M mirroring upstream slothe.cpp (ggml) math: F16 islands,
ternary linears = SubLN -> (A8: per-256-block absmax int8, ggml Q8_K) -> {-1,0,1}*scale.
Used to check the Core ML graph structure and activation ranges; ggml itself is the reference."""
import json
import numpy as np
W = np.load("work/w12m.npz"); C = json.load(open("work/cfg12m.json"))
EPS = 1e-6
f16 = lambda a: a.astype(np.float16).astype(np.float32)
def rms(x, w):
return x / np.sqrt((x * x).mean(-1, keepdims=True) + EPS) * f16(w)
def a8(x):
L, n = x.shape
xb = x.reshape(L, n // 256, 256)
am = np.abs(xb).max(-1, keepdims=True)
am = np.where(am == 0, 1, am)
return (np.round(xb * (127.0 / am)) * (am / 127.0)).reshape(L, n)
STATS = {}
TRACE = []
def lin(i, s, x, fp, act8=True):
if fp:
return f16(x) @ f16(W[f"{i}.{s}.w"]).T
x = rms(x, W[f"{i}.{s}.pre"])
STATS["lin_in"] = max(STATS.get("lin_in", 0), float(np.abs(x).max()))
if act8:
x = a8(x)
return (x @ W[f"{i}.{s}.q"].astype(np.float32).T) * f16(W[f"{i}.{s}.s"])
def rope_tab(L, hd=32):
half = hd // 2
freq = 1.0 / (10000 ** (np.arange(half) / half))
ang = np.arange(L)[:, None] * freq[None]
return np.concatenate([np.cos(ang)] * 2, -1), np.concatenate([np.sin(ang)] * 2, -1)
def forward(ids, act8=True):
L = len(ids); D = C["depth"]; H, KV, hd = C["heads"], C["kv"], C["dim"] // C["heads"]
cos, sin = rope_tab(L, hd); TRACE.clear()
x = rms(f16(W["embed"])[ids], W["embed_norm"])
for i in range(D):
fp = i == 0 or i == D - 1
h = rms(x, W[f"{i}.n1"])
q = lin(i, "q", h, fp, act8).reshape(L, H, hd).transpose(1, 0, 2)
k = lin(i, "k", h, fp, act8).reshape(L, KV, hd).transpose(1, 0, 2)
v = lin(i, "v", h, fp, act8).reshape(L, KV, hd).transpose(1, 0, 2)
q, k = rms(q, W[f"{i}.qn"]), rms(k, W[f"{i}.kn"])
rot = lambda t: np.concatenate([-t[..., hd // 2:], t[..., :hd // 2]], -1)
q, k = q * cos + rot(q) * sin, k * cos + rot(k) * sin
k, v = np.repeat(k, H // KV, 0), np.repeat(v, H // KV, 0)
s = q @ k.transpose(0, 2, 1) / np.sqrt(hd)
s = np.exp(s - s.max(-1, keepdims=True)); s /= s.sum(-1, keepdims=True)
o = (s @ v).transpose(1, 0, 2).reshape(L, H * hd)
x = x + lin(i, "o", o, fp, act8)
h2 = rms(x, W[f"{i}.n2"])
g, u = lin(i, "w1", h2, fp, act8), lin(i, "w3", h2, fp, act8)
ff = g / (1 + np.exp(-g)) * u
STATS["ffn_act"] = max(STATS.get("ffn_act", 0), float(np.abs(ff).max()))
x = x + lin(i, "w2", ff, fp, act8)
STATS["resid"] = max(STATS.get("resid", 0), float(np.abs(x).max()))
TRACE.append(x.copy())
x = rms(x, W["norm"])
return f16(x) @ f16(W["head"]).T