"""Task 2: convert QDec to Core ML. usage: convert.py B T [variant...]; variants fp16 | pal8 | lin8. Writes models/dec_b{B}_t{T}_{variant}.mlpackage and appends one line per variant to logs/convert.log.""" import sys, time, traceback, torch, numpy as np, coremltools as ct import coremltools.optimize.coreml as cto from pathlib import Path from qdec import QDec, ane_kwargs B, T = int(sys.argv[1]), int(sys.argv[2]) variants = sys.argv[3:] or ["fp16"] Path("models").mkdir(exist_ok=True); Path("logs").mkdir(exist_ok=True) log = open("logs/convert.log", "a") def L(msg): print(msg); log.write(msg + "\n"); log.flush() import os CH = int(os.environ.get('CHUNK', '0')) or None m = QDec(T, B=B, chunk=CH, **ane_kwargs()).eval() SUF = f'c{CH}' if CH else '' ex = torch.zeros((B, T), dtype=torch.int32) t0 = time.time() with torch.no_grad(): tr = torch.jit.trace(m, ex.long()) base = None for v in variants: name = f"models/dec_b{B}_t{T}_{v}{SUF}.mlpackage" try: t0 = time.time() if base is None: base = ct.convert(tr, inputs=[ct.TensorType("ids", shape=(B, 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") mm = base if v == "pal8": mm = cto.palettize_weights(base, cto.OptimizationConfig(global_config=cto.OpPalettizerConfig(nbits=8, mode="kmeans"))) elif v == "lin8": mm = cto.linear_quantize_weights(base, cto.OptimizationConfig(global_config=cto.OpLinearQuantizerConfig(mode="linear_symmetric"))) mm.save(name) sz = sum(f.stat().st_size for f in Path(name).rglob("*") if f.is_file()) / 1e6 L(f"OK b{B} t{T} {v}{SUF}: {time.time()-t0:.1f}s {sz:.1f} MB -> {name}") except Exception as e: L(f"FAIL b{B} t{T} {v}{SUF}: {type(e).__name__}: {str(e).splitlines()[0][:300]}") traceback.print_exc()