McBopomofoLM-models / scripts /enc /build_coreml.py
workfunction's picture
McBopomofoLM v2.1.1 Core ML models (SlothE-T 25M encoder, pred_q35_60m decoder)
ff5f59d
Raw
History Blame Contribute Delete
10.1 kB
"""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)