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