| """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: |
| 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: |
| 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 = 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): |
| 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] |
| SIN = h16(np.concatenate([np.sin(ang)] * 2, -1))[None] |
| SIGN = h16(np.concatenate([-np.ones(half), np.ones(half)])) |
|
|
| def rope(t): |
| 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: |
| 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) |
| 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 |
| 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) |
| 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"])) |
| 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 |
| 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: |
| 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) |
|
|