import sys, time import numpy as np, torch, coremltools as ct from eos_graph import load_chunked, load_eos, load_rows # packed: convert_eos.py S P N | rows: convert_eos.py S BxQ N S, N = int(sys.argv[1]), int(sys.argv[3]) if "x" in sys.argv[2]: B, Q = map(int, sys.argv[2].split("x")); P = B * Q g = load_rows("eos", S, B, Q); tag = f"S{S}_R{B}x{Q}_N{N}" elif sys.argv[2].startswith("c"): P = int(sys.argv[2][1:]); g = load_chunked("eos", S, P); tag = f"S{S}_C{P}_N{N}" else: P = int(sys.argv[2]); g = load_eos("eos", S, P); tag = f"S{S}_P{P}_N{N}" chunked = sys.argv[2].startswith("c") out = sys.argv[4] if len(sys.argv) > 4 else f"coreml/eos_{tag}.mlpackage" T, R = S + P, g.cfg.rotary_dim ex = (torch.zeros(1, T, dtype=torch.int32), torch.zeros(T, R), torch.zeros(T, R), torch.ones(S), torch.zeros(3, S), torch.eye(P), torch.zeros(3, P), torch.zeros(3, P, 3)) if chunked: M = P // 64 ex += (torch.eye(64).expand(M, 64, 64).contiguous(), torch.zeros(P), torch.zeros(M, 64)) ex += (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, check_trace=False) f16 = np.float16 m = ct.convert(tr, convert_to="mlprogram", minimum_deployment_target=ct.target.macOS15, inputs=[ct.TensorType(name="input_ids", shape=(1, T), dtype=np.int32), ct.TensorType(name="cos", shape=(T, R), dtype=f16), ct.TensorType(name="sin", shape=(T, R), dtype=f16), ct.TensorType(name="valid", shape=(S,), dtype=f16), ct.TensorType(name="tail_onehot", shape=(3, S), dtype=f16), ct.TensorType(name="segment", shape=(P, P), dtype=f16), ct.TensorType(name="lag_keep", shape=(3, P), dtype=f16), ct.TensorType(name="lag_tail", shape=(3, P, 3), dtype=f16)] + ([ct.TensorType(name="seg_chunks", shape=(P // 64, 64, 64), dtype=f16), ct.TensorType(name="cont", shape=(P,), dtype=f16), ct.TensorType(name="last_seg", shape=(P // 64, 64), dtype=f16)] if chunked else []) + [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-Eos-0.8B (vllm-sr): shared prefix + packed questions, one call" m.save(out); print("saved", out, f"{time.time()-t:.0f}s")