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