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
4.39 kB
"""Host side: System One request -> packed tree inputs, and logits -> answers (vendored code)."""
import sys
import numpy as np
sys.path.insert(0, "kai")
from decision2._vendor.dev2model.decision_model import encode # noqa: E402
from decision2._vendor.dev2model.infer import product_answer, question_to_row # noqa: E402
from decision2._vendor.dev2model.score_bias import apply as apply_score_bias # noqa: E402
NEG = -1e4
def rows(tokenizer, state, questions, cap=8192):
item = {"id": "request", "state": state}
out = []
for qid, q in questions.items():
row = question_to_row(item, qid, q)
out.append((qid, row, encode(row, tokenizer, cap)))
return out
def pack(jobs, L, N, pad_id):
"""Shared prefix once, then each suffix; suffix tokens see prefix + own suffix (causal)."""
seqs = [e["ids"] for _, _, e in jobs]
P = min(min(e["candidate_positions"][0] for _, _, e in jobs), min(len(s) for s in seqs) - 1)
for i in range(P):
if any(s[i] != seqs[0][i] for s in seqs):
P = i
break
ids, pos, seg = list(seqs[0][:P]), list(range(P)), [-1] * P
starts = []
for j, s in enumerate(seqs):
starts.append(len(ids) - P)
ids += s[P:]
pos += range(P, len(s))
seg += [j] * (len(s) - P)
T = len(ids)
if T > L:
raise ValueError(f"packed length {T} > {L}")
cand, qry, owner = [], [], []
for j, (_, _, e) in enumerate(jobs):
for c in e["candidate_positions"]:
cand.append(c if c < P else c + starts[j])
qry.append(e["query_position"] + starts[j])
owner.append(j)
if len(cand) > N:
raise ValueError(f"{len(cand)} candidates > {N}")
seg = np.array(seg + [-2] * (L - T))
p = np.array(pos + [0] * (L - T))
i = np.arange(L)
causal = i[None, :] <= i[:, None]
same = (seg[None, :] == seg[:, None]) | (seg[None, :] == -1)
allow = causal & same & (seg[None, :] != -2)
allow[np.arange(L), np.arange(L)] = True # padding rows attend to themselves
mask = np.where(allow, 0.0, NEG).astype(np.float32)[None, None]
pad = N - len(cand)
return {
"input_ids": np.array([ids + [pad_id] * (L - T)], dtype=np.int32),
"position_ids": p[None].astype(np.int32),
"mask": mask,
"cand_idx": np.array(cand + [0] * pad, dtype=np.int32),
"query_idx": np.array(qry + [0] * pad, dtype=np.int32),
}, owner, T
def answers(jobs, logits, owner, score_bias=None, temps=None):
per = [[] for _ in jobs]
for j, v in zip(owner, logits[: len(owner)]):
per[j].append(float(v))
out = {}
for (qid, row, e), values in zip(jobs, per):
if score_bias is not None and row["task_type"] == "score":
values = apply_score_bias(score_bias, values, len(row["options"]))
out[qid] = product_answer(
row["task_type"], e["keys"], values, (temps or {}).get(row["task_type"], 1.0),
[o["description"] for o in row["options"]],
)
return out
def size(jobs):
"""(packed tokens, candidates) of one packed call, without building it."""
seqs = [e["ids"] for _, _, e in jobs]
P = min(min(e["candidate_positions"][0] for _, _, e in jobs), min(len(s) for s in seqs) - 1)
for i in range(P):
if any(s[i] != seqs[0][i] for s in seqs):
P = i
break
return P + sum(len(s) - P for s in seqs), sum(len(e["keys"]) for _, _, e in jobs)
def fits(jobs, L, N):
T, C = size(jobs)
return T <= L and C <= N
def chunks(jobs, L, N):
"""Greedy split of a request's questions into groups that each pack into one L/N call."""
groups, cur = [], []
for j in range(len(jobs)):
if not cur or fits([jobs[i] for i in cur + [j]], L, N):
cur.append(j)
else:
groups.append(cur)
cur = [j]
return groups + [cur]
def run(model, jobs, L, N, pad_id):
"""Logits per job (list of lists), over as many packed calls as needed."""
per = [None] * len(jobs)
for group in chunks(jobs, L, N):
sub = [jobs[i] for i in group]
x, owner, _ = pack(sub, L, N, pad_id)
x["mask"] = x["mask"].astype("float16")
out = model.predict(x)["logits"]
for k, i in enumerate(group):
per[i] = [float(v) for v, o in zip(out, owner) if o == k]
return per