import sys, time import numpy as np, torch, coremltools as ct from kai_graph import KaiGraph L, N = int(sys.argv[1]), int(sys.argv[2]) out = sys.argv[3] if len(sys.argv) > 3 else f"coreml/kai_L{L}_N{N}.mlpackage" g = KaiGraph().eval() ex = (torch.zeros(1, L, dtype=torch.int32), torch.arange(L, dtype=torch.int32)[None], torch.zeros(1, 1, L, L), torch.zeros(N, dtype=torch.int32), torch.zeros(N, dtype=torch.int32)) t = time.time() with torch.no_grad(): tr = torch.jit.trace(g, ex) m = ct.convert(tr, convert_to="mlprogram", minimum_deployment_target=ct.target.macOS15, inputs=[ct.TensorType(name="input_ids", shape=(1, L), dtype=np.int32), ct.TensorType(name="position_ids", shape=(1, L), dtype=np.int32), ct.TensorType(name="mask", shape=(1, 1, L, L), dtype=np.float16), ct.TensorType(name="cand_idx", shape=(N,), dtype=np.int32), ct.TensorType(name="query_idx", shape=(N,), dtype=np.int32)], outputs=[ct.TensorType(name="logits", dtype=np.float32)], compute_precision=ct.precision.FLOAT16) m.short_description = "Decision-2.0-Kai-0.6B (vllm-sr) packed multi-question decision graph" m.save(out); print("saved", out, f"{time.time()-t:.0f}s")