"""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 /parts/dec_b3_t{T}_{v}.mlpackage, /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]}")