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