workfunction's picture
McBopomofoLM v2.1.1 Core ML models (SlothE-T 25M encoder, pred_q35_60m decoder)
ff5f59d
Raw
History Blame Contribute Delete
2.7 kB
"""Phase F: build the shipped decoder = ONE multifunction Core ML model (functions t16/t32/t64/t96, B=3, shared
weights; t16 single-block GDN, t>=32 chunked c16), fp16 and pal8. usage: build_mf.py [outdir=mf]
Writes <outdir>/parts/dec_b3_t{T}_{v}.mlpackage, <outdir>/dec_mf_{v}.mlpackage, appends to logs/convert.log."""
import sys, time, shutil, torch, numpy as np, coremltools as ct
import coremltools.optimize.coreml as cto
from coremltools.models.utils import MultiFunctionDescriptor, save_multifunction
from pathlib import Path
from qdec import QDec, ane_kwargs
out = Path(sys.argv[1] if len(sys.argv) > 1 else "mf"); (out / "parts").mkdir(parents=True, exist_ok=True)
log = open("logs/convert.log", "a")
def L(m): print(m, flush=True); log.write(m + "\n"); log.flush()
def size(p): return sum(f.stat().st_size for f in Path(p).rglob("*") if f.is_file()) / 1e6
L(f"--- Phase F multifunction build {time.ctime()}")
parts = {"fp16": {}, "pal8": {}}
for T in (16, 32, 64, 96):
m = QDec(T, B=3, chunk=16 if T > 16 else None, **ane_kwargs()).eval()
with torch.no_grad(): tr = torch.jit.trace(m, torch.zeros((3, T), dtype=torch.long))
t0 = time.time()
mm = ct.convert(tr, inputs=[ct.TensorType("ids", shape=(3, T), dtype=np.int32)], outputs=[ct.TensorType("lp", dtype=np.float32)],
compute_precision=ct.precision.FLOAT16, minimum_deployment_target=ct.target.macOS15, convert_to="mlprogram")
p = out / "parts" / f"dec_b3_t{T}_fp16.mlpackage"; mm.save(str(p)); parts["fp16"][T] = p
L(f"OK part t{T} fp16 {time.time()-t0:.1f}s {size(p):.1f} MB")
for v in ("fp16",):
d = MultiFunctionDescriptor()
for T, p in parts[v].items(): d.add_function(str(p), src_function_name="main", target_function_name=f"t{T}")
d.default_function_name = "t16"
dst = out / f"dec_mf_{v}.mlpackage"; shutil.rmtree(dst, ignore_errors=True); save_multifunction(d, str(dst))
L(f"OK multifunction {v}: {size(dst):.1f} MB (separate parts total {sum(size(p) for p in parts[v].values()):.1f} MB) -> {dst}")
# pal8: palettize the merged fp16 multifunction model (one k-means per shared weight -> still shared)
t0 = time.time()
try:
mf = ct.models.MLModel(str(out / "dec_mf_fp16.mlpackage"), skip_model_load=True)
pm = cto.palettize_weights(mf, cto.OptimizationConfig(global_config=cto.OpPalettizerConfig(nbits=8, mode="kmeans")))
dst = out / "dec_mf_pal8.mlpackage"; shutil.rmtree(dst, ignore_errors=True); pm.save(str(dst))
L(f"OK multifunction pal8 (palettized after merge): {size(dst):.1f} MB, {time.time()-t0:.0f}s -> {dst}")
except Exception as e:
L(f"FAIL multifunction pal8 after merge: {type(e).__name__}: {str(e).splitlines()[0][:300]}")