File size: 2,777 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
65
"""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