alexwengg's picture
Decision-2.0-Eos-0.8B Core ML: shared prefix + chunked packed questions, fp16 package, runtime, parity reports
7bf8323 verified
Raw History Blame Contribute Delete
2.4 kB
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")