File size: 10,051 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
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
"""Build SlothE-T 12M Core ML ML Programs directly with the MIL builder (no torch tracing).

Math mirrors upstream slothe.cpp (ggml): F16 islands (embed, blocks 0/11, head), ternary linears as
SubLN -> [A8 fake-quant per 256-block, ggml Q8_K] -> x @ {-1,0,1}^T * scale[out]  (int-dot-then-scale,
like ggml_vec_dot_tq2_0_q8_K), GQA 8q/2kv, QK-norm, NEOX RoPE, SwiGLU, no causal mask.
Fixed-length buckets: input ids int32 [1,L] (+ key mask fp16 [1,L], 1=real 0=pad); padded keys get
-1e4 so real positions are exactly the unpadded forward. Output logits fp16 [1,L,8342]
(legal-mask log-softmax on CPU) or, with --lsm, masked log-probs from an in-graph legal table.

usage: build_coreml.py --L 16 --wmode fp16|pal2|w8 [--act8] [--lsm] [--out path.mlpackage]
"""
import argparse, json
import numpy as np
import coremltools as ct
from coremltools.converters.mil import Builder as mb
from coremltools.converters.mil.mil import types

ap = argparse.ArgumentParser()
ap.add_argument("--L", type=int, default=16)
ap.add_argument("--size", default="12m", choices=["12m", "25m"])
ap.add_argument("--wmode", default="fp16", choices=["fp16", "pal2", "w8"])
ap.add_argument("--act8", action="store_true")
ap.add_argument("--lsm", action="store_true", help="in-graph legal-mask log-softmax")
ap.add_argument("--out")
ap.add_argument("--taps", action="store_true", help="debug taps inside block 1")
ap.add_argument("--prescale", action="store_true", help="RMSNorm on x/absmax(x): avoids fp16 subnormal x^2 flush on ANE")
ap.add_argument("--silu", default="tanh", choices=["silu", "sigmoid", "tanh", "exp"], help="SiLU formulation (ANE silu op is approximate)")
ap.add_argument("--newton", type=int, default=0, help="Newton steps refining rsqrt (ANE rsqrt is approximate)")
ap.add_argument("--embin", action="store_true", help="input = embedding rows fp16 [1,L,256] gathered by the caller (keeps the whole graph on ANE)")
ap.add_argument("--salt", type=int, default=0, help="perturb the pad-key bias constant (1e4 + 8*salt; no effect on outputs) -> a never-seen model hash for cold-load timing")
ap.add_argument("--debug", action="store_true", help="also output every block residual")
ap.add_argument("--rmsln", action="store_true", help="RMSNorm via layer_norm on [x,-x]")
a = ap.parse_args()
W = np.load(f"work/w{a.size}.npz"); C = json.load(open(f"work/cfg{a.size}.json"))
L, D, H, KV, DIM = a.L, C["depth"], C["heads"], C["kv"], C["dim"]
HD = DIM // H; EPS = 1e-6
h16 = lambda x: np.asarray(x, dtype=np.float16)

def rms(x, w, axes=(-1,)):
    if a.rmsln:   # mean([x,-x]) = 0 -> layer_norm == rms_norm; ANE has a native, better-conditioned layer_norm
        n = x.shape[-1]
        y = mb.layer_norm(x=mb.concat(values=[x, mb.mul(x=x, y=h16(-1.0))], axis=-1), axes=[-1], epsilon=EPS)
        y = mb.slice_by_index(x=y, begin=[0] * (len(x.shape) - 1) + [0], end=list(x.shape[:-1]) + [n])
        return mb.mul(x=y, y=h16(w))
    if a.prescale:   # rms(x) == rms(x/c); c = row absmax keeps x^2 in fp16 normal range (ANE flushes subnormals)
        inv = mb.real_div(x=h16(1.0), y=mb.maximum(x=mb.reduce_max(x=mb.abs(x=x), axes=list(axes), keep_dims=True), y=h16(1.0 / 256)))
        xs = mb.mul(x=x, y=inv)
        ms = mb.reduce_mean(x=mb.mul(x=xs, y=xs), axes=list(axes), keep_dims=True)
        ms = mb.add(x=ms, y=mb.mul(x=mb.mul(x=inv, y=inv), y=h16(EPS)))
        return mb.mul(x=mb.mul(x=xs, y=mb.rsqrt(x=ms)), y=h16(w))
    ms = mb.reduce_mean(x=mb.mul(x=x, y=x), axes=list(axes), keep_dims=True)
    ms = mb.add(x=ms, y=h16(EPS)) if a.newton else ms
    r = mb.rsqrt(x=ms, epsilon=EPS) if not a.newton else mb.rsqrt(x=ms)
    for _ in range(a.newton):   # r <- r * (1.5 - 0.5 * ms * r^2)
        r = mb.mul(x=r, y=mb.sub(x=h16(1.5), y=mb.mul(x=mb.mul(x=ms, y=h16(0.5)), y=mb.mul(x=r, y=r))))
    return mb.mul(x=mb.mul(x=x, y=r), y=h16(w))

def act8(x, n):   # x [L, n] -> per-256-block absmax int8 fake-quant (ggml Q8_K math, iscale=127/amax)
    xb = mb.reshape(x=x, shape=[L, n // 256, 256])
    am = mb.reduce_max(x=mb.abs(x=xb), axes=[-1], keep_dims=True)
    am = mb.maximum(x=am, y=h16(1e-5))
    q = mb.round(x=mb.mul(x=xb, y=mb.real_div(x=h16(127.0), y=am)))
    return mb.reshape(x=mb.mul(x=q, y=mb.real_div(x=am, y=h16(127.0))), shape=[L, n])

def silu(g):
    if a.silu == "silu": return mb.silu(x=g)
    if a.silu == "sigmoid": return mb.mul(x=g, y=mb.sigmoid(x=g))
    if a.silu == "tanh": return mb.mul(x=g, y=mb.add(x=h16(0.5), y=mb.mul(x=h16(0.5), y=mb.tanh(x=mb.mul(x=g, y=h16(0.5))))))
    return mb.real_div(x=g, y=mb.add(x=h16(1.0), y=mb.exp(x=mb.mul(x=g, y=h16(-1.0)))))

TAPS = []
def tap(i, name, v):
    if a.taps and i == 1:
        TAPS.append(mb.identity(x=v, name=f't_{name}'))
    return v

def lin(i, s, x, fp, n_in):
    if fp:
        return mb.linear(x=x, weight=h16(W[f"{i}.{s}.w"]))
    x = tap(i, f"pre_{s}", rms(x, W[f"{i}.{s}.pre"]))
    if a.act8:
        x = act8(x, n_in)
    y = mb.linear(x=x, weight=h16(W[f"{i}.{s}.q"]), name=f"tern_{i}_{s}")
    tap(i, f"raw_{s}", y)
    return tap(i, f"lin_{s}", mb.mul(x=y, y=h16(W[f"{i}.{s}.s"])))

half = HD // 2
freq = 1.0 / (10000 ** (np.arange(half) / half))
ang = np.arange(L)[:, None] * freq[None]
COS = h16(np.concatenate([np.cos(ang)] * 2, -1))[None]   # [1, L, HD]
SIN = h16(np.concatenate([np.sin(ang)] * 2, -1))[None]
SIGN = h16(np.concatenate([-np.ones(half), np.ones(half)]))  # rot(x) = [-x2, x1]

def rope(t):   # t [nh, L, HD]
    x1 = mb.slice_by_index(x=t, begin=[0, 0, 0], end=[0, 0, half], end_mask=[True, True, False])
    x2 = mb.slice_by_index(x=t, begin=[0, 0, half], end=[0, 0, 0], end_mask=[True, True, True])
    rot = mb.mul(x=mb.concat(values=[x2, x1], axis=-1), y=SIGN)
    return mb.add(x=mb.mul(x=t, y=COS), y=mb.mul(x=rot, y=SIN))

inputs = {"ids": mb.TensorSpec(shape=(1, L, DIM) if a.embin else (1, L), dtype=types.fp16 if a.embin else types.int32),
          "mask": mb.TensorSpec(shape=(1, L), dtype=types.fp16)}
if a.lsm:
    LEGAL_BIAS = h16(np.where(np.load(__import__("common").hf("syl2legal.npz"))["mask"], 0.0, -1e4))

@mb.program(input_specs=list(inputs.values()), opset_version=ct.target.macOS15)
def prog(ids, mask):
    if a.embin:   # 'ids' carries the caller-gathered fp16 embedding rows
        x = mb.reshape(x=ids, shape=[L, DIM])
    else:
        idv = mb.reshape(x=ids, shape=[L])
        x = mb.gather(x=h16(W["embed"]), indices=idv, axis=0)        # [L, DIM]
    x = rms(x, W["embed_norm"])
    kb = mb.reshape(x=mb.mul(x=mb.sub(x=mask, y=h16(1.0)), y=h16(1e4 + 8 * a.salt)), shape=[1, 1, L])
    dbg = []
    for i in range(D):
        fp = i == 0 or i == D - 1
        h = tap(i, "h", rms(x, W[f"{i}.n1"]))
        q = mb.transpose(x=mb.reshape(x=lin(i, "q", h, fp, DIM), shape=[L, H, HD]), perm=[1, 0, 2])
        k = mb.transpose(x=mb.reshape(x=lin(i, "k", h, fp, DIM), shape=[L, KV, HD]), perm=[1, 0, 2])
        v = mb.transpose(x=mb.reshape(x=lin(i, "v", h, fp, DIM), shape=[L, KV, HD]), perm=[1, 0, 2])
        q = tap(i, "qr", rope(tap(i, "qn", rms(q, W[f"{i}.qn"])))); k = rope(rms(k, W[f"{i}.kn"]))
        rep = H // KV   # repeat_interleave on head axis: head j uses kv j // rep
        k = mb.reshape(x=mb.tile(x=mb.reshape(x=k, shape=[KV, 1, L, HD]), reps=[1, rep, 1, 1]), shape=[H, L, HD])
        v = mb.reshape(x=mb.tile(x=mb.reshape(x=v, shape=[KV, 1, L, HD]), reps=[1, rep, 1, 1]), shape=[H, L, HD])
        s = mb.matmul(x=mb.mul(x=q, y=h16(1.0 / np.sqrt(HD))), y=k, transpose_y=True)   # [H, L, L]
        p = tap(i, "probs", mb.softmax(x=mb.add(x=tap(i, "scores", s), y=kb), axis=-1))
        o = mb.reshape(x=mb.transpose(x=mb.matmul(x=p, y=v), perm=[1, 0, 2]), shape=[L, H * HD])
        x = tap(i, "x_attn", mb.add(x=x, y=lin(i, "o", tap(i, "o", o), fp, DIM)))
        h2 = rms(x, W[f"{i}.n2"])
        g = lin(i, "w1", h2, fp, DIM); u = lin(i, "w3", h2, fp, DIM)
        x = mb.add(x=x, y=lin(i, "w2", tap(i, "ff", mb.mul(x=tap(i, "silu", silu(g)), y=u)), fp, C["ffn"]))
        if a.debug:
            dbg.append(mb.identity(x=x, name=f"blk{i}"))
    x = rms(x, W["norm"])
    lg = mb.linear(x=x, weight=h16(W["head"]))                          # [L, n_char]
    if a.lsm:
        lg = mb.add(x=lg, y=mb.gather(x=LEGAL_BIAS, indices=idv, axis=0))
        lg = mb.sub(x=lg, y=mb.reduce_log_sum_exp(x=lg, axes=[-1], keep_dims=True))
    out = mb.reshape(x=lg, shape=[1, L, C["n_char"]], name="logits")
    return (out, *dbg, *TAPS) if (a.debug or a.taps) else out

m = ct.convert(prog, convert_to="mlprogram", compute_precision=ct.precision.FLOAT16,
               minimum_deployment_target=ct.target.macOS15, compute_units=ct.ComputeUnit.ALL,
               skip_model_load=True)
if a.wmode in ("pal2", "w8"):
    from coremltools.converters.mil.mil.passes.graph_pass import PassOption  # noqa
    import coremltools.optimize.coreml as cto
    tern_consts = set()
    for f in m._mil_program.functions.values():
        for op in f.operations:
            if op.op_type == "linear" and op.name.startswith("tern_"):
                tern_consts.add(op.weight.op.name)
    if a.wmode == "pal2":
        cfg = cto.OptimizationConfig(op_name_configs={n: cto.OpPalettizerConfig(mode="unique", weight_threshold=1)
                                                      for n in tern_consts})
        m = cto.palettize_weights(m, cfg)
    else:   # int8 per-channel symmetric on all weights >= 2048 elements (ternary codes exact in int8)
        cfg = cto.OptimizationConfig(global_config=cto.OpLinearQuantizerConfig(mode="linear_symmetric", dtype="int8",
                                                                               granularity="per_channel", weight_threshold=2048))
        m = cto.linear_quantize_weights(m, cfg)
out = a.out or f"models/enc{a.size}_L{L}_{a.wmode}{'_a8' if a.act8 else ''}{'_lsm' if a.lsm else ''}{'_ln' if a.rmsln else ''}{'_emb' if a.embin else ''}{'_ps' if a.prescale else ''}{'_' + a.silu if a.silu != 'tanh' else ''}{'_dbg' if a.debug else ''}{'_taps' if a.taps else ''}{'_nt%d' % a.newton if a.newton else ''}.mlpackage"
m.save(out)
print("saved", out)