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