"""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