File size: 4,392 Bytes
7bf8323
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
"""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