"""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