File size: 8,828 Bytes
03223d7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
from __future__ import annotations
import json
import torch


def encode_piece(tok, text: str) -> list[int]:
    return tok(text, add_special_tokens=False).input_ids


def state_to_text(state) -> str:
    if isinstance(state, str):
        return state
    return json.dumps(state, ensure_ascii=False, sort_keys=True, separators=(",", ":"))


def _clip_head_tail(ids: list[int], limit: int) -> list[int]:
    """Deterministically keep evidence from both ends of a long field."""
    if limit <= 0:
        return []
    if len(ids) <= limit:
        return ids
    head = (limit + 1) // 2
    tail = limit - head
    return ids[:head] + (ids[-tail:] if tail else [])


def _assemble(tok, state_ids, q_ids, option_ids, max_length: int):
    ids = [tok.bos_token_id] if tok.bos_token_id is not None else []
    ids += encode_piece(tok, "State:\n")
    ids += state_ids
    ids += encode_piece(tok, "\n\nQuestion:\n")
    ids += q_ids
    ids += encode_piece(tok, "\n\nOptions:\n")

    option_positions = []
    option_spans = []
    for idx, opt_ids in enumerate(option_ids):
        ids += encode_piece(tok, f"- [{idx}] ")
        span_start = len(ids)
        ids += opt_ids
        span_end = len(ids)
        option_spans.append((span_start, span_end))
        # Represent each option by its final semantic token, after it has seen
        # the full option text but before the newline delimiter.
        option_positions.append(len(ids) - 1)
        ids += encode_piece(tok, "\n")

    ids += encode_piece(tok, "\nDecision:")
    decide_position = len(ids) - 1
    if len(ids) > max_length:
        return None
    return ids, option_positions, decide_position, option_spans


def pack_question(tok, state, question: dict, max_length: int):
    """Pack a decision example without dropping long states.

    Priority order is structural markers/options/question first, then state.
    Long state is head+tail truncated. If question/options themselves are huge,
    they are bounded as a second-stage fallback. Returning None is reserved for
    structurally impossible cases (e.g. too many options for max_length).
    """
    if max_length < 32:
        return None

    state_ids = encode_piece(tok, state_to_text(state))
    q_ids = encode_piece(tok, str(question["instruction"]))
    option_ids = [encode_piece(tok, str(opt["text"])) for opt in question["options"]]
    if not option_ids or any(not x for x in option_ids):
        return None

    was_truncated = False

    # First discover how much room remains for the state while preserving the
    # complete question and option text.
    probe = _assemble(tok, [], q_ids, option_ids, max_length=10**9)
    if probe is None:
        return None
    structural_len = len(probe[0])
    if structural_len <= max_length:
        state_budget = max_length - structural_len
        was_truncated = len(state_ids) > state_budget
        packed = _assemble(tok, _clip_head_tail(state_ids, state_budget), q_ids, option_ids, max_length)
    else:
        was_truncated = True
        # Extremely verbose question/tool schemas: cap semantic fields rather
        # than dropping the sample. Typical training examples never hit this.
        nopt = len(option_ids)
        q_cap = min(len(q_ids), max(16, max_length // 8))
        # Start modestly; then shrink until the structural representation fits.
        opt_cap = max(8, min(96, max_length // max(8, nopt * 2)))
        q_fit = _clip_head_tail(q_ids, q_cap)
        opts_fit = [_clip_head_tail(x, opt_cap) for x in option_ids]
        packed = _assemble(tok, [], q_fit, opts_fit, max_length)
        while packed is None and (q_cap > 8 or opt_cap > 4):
            q_cap = max(8, q_cap // 2)
            opt_cap = max(4, opt_cap // 2)
            q_fit = _clip_head_tail(q_ids, q_cap)
            opts_fit = [_clip_head_tail(x, opt_cap) for x in option_ids]
            packed = _assemble(tok, [], q_fit, opts_fit, max_length)

        if packed is not None:
            # If the shrunken question/options leave room, fill it with state.
            base_len = len(packed[0])
            state_budget = max(0, max_length - base_len)
            was_truncated = was_truncated or len(state_ids) > state_budget
            packed = _assemble(tok, _clip_head_tail(state_ids, state_budget), q_fit, opts_fit, max_length)

    if packed is None:
        return None
    ids, option_positions, decide_position, option_spans = packed
    return {
        "input_ids": ids,
        "option_positions": option_positions,
        "option_spans": option_spans,
        "decide_position": decide_position,
        "target": question.get("target_distribution"),
        "was_truncated": was_truncated,
    }


def collate_packed(tok, packed: list[dict]) -> dict[str, torch.Tensor]:
    if not packed:
        raise RuntimeError("No packable examples")
    bsz = len(packed)
    max_seq = max(len(x["input_ids"]) for x in packed)
    max_opts = max(len(x["option_positions"]) for x in packed)
    pad_id = tok.pad_token_id if tok.pad_token_id is not None else 0
    input_ids = torch.full((bsz, max_seq), pad_id, dtype=torch.long)
    attention_mask = torch.zeros((bsz, max_seq), dtype=torch.long)
    option_positions = torch.zeros((bsz, max_opts), dtype=torch.long)
    option_mask = torch.zeros((bsz, max_opts), dtype=torch.bool)
    decide_positions = torch.zeros((bsz,), dtype=torch.long)
    targets = torch.zeros((bsz, max_opts), dtype=torch.float32)
    option_starts = torch.zeros((bsz, max_opts), dtype=torch.long)
    option_ends = torch.zeros((bsz, max_opts), dtype=torch.long)
    has_targets = all(x.get("target") is not None for x in packed)
    for i, x in enumerate(packed):
        n = len(x["input_ids"])
        m = len(x["option_positions"])
        input_ids[i, :n] = torch.tensor(x["input_ids"], dtype=torch.long)
        attention_mask[i, :n] = 1
        option_positions[i, :m] = torch.tensor(x["option_positions"], dtype=torch.long)
        spans = x.get("option_spans") or [(int(v), int(v)+1) for v in x["option_positions"]]
        option_starts[i, :m] = torch.tensor([a for a, _ in spans], dtype=torch.long)
        option_ends[i, :m] = torch.tensor([b for _, b in spans], dtype=torch.long)
        option_mask[i, :m] = True
        decide_positions[i] = x["decide_position"]
        if has_targets:
            targets[i, :m] = torch.tensor(x["target"], dtype=torch.float32)
    out = {
        "input_ids": input_ids,
        "attention_mask": attention_mask,
        "option_positions": option_positions,
        "option_starts": option_starts,
        "option_ends": option_ends,
        "option_mask": option_mask,
        "decide_positions": decide_positions,
    }
    if has_targets:
        out["targets"] = targets
    return out


def pack_shared_request(tok, state, questions: list[dict], max_length: int):
    """Pack one common state prefix plus causal question suffixes.

    Every question sees exactly the same serialized/clipped state. This is the
    representation needed for safe KV-prefix reuse at inference time.
    """
    if max_length < 32 or not questions:
        return None
    bos = [tok.bos_token_id] if tok.bos_token_id is not None else []
    state_marker = encode_piece(tok, "State:\n")
    state_ids = encode_piece(tok, state_to_text(state))

    suffixes = []
    max_suffix_len = 0
    for question in questions:
        q_ids = encode_piece(tok, str(question["instruction"]))
        option_ids = [encode_piece(tok, str(opt["text"])) for opt in question["options"]]
        if not option_ids or any(not x for x in option_ids):
            return None
        ids = encode_piece(tok, "\n\nQuestion:\n") + q_ids + encode_piece(tok, "\n\nOptions:\n")
        option_positions = []
        for idx, opt_ids in enumerate(option_ids):
            ids += encode_piece(tok, f"- [{idx}] ")
            ids += opt_ids
            option_positions.append(len(ids) - 1)
            ids += encode_piece(tok, "\n")
        ids += encode_piece(tok, "\nDecision:")
        decide_position = len(ids) - 1
        suffixes.append({
            "input_ids": ids,
            "option_positions": option_positions,
            "decide_position": decide_position,
        })
        max_suffix_len = max(max_suffix_len, len(ids))

    fixed_prefix_len = len(bos) + len(state_marker)
    state_budget = max_length - fixed_prefix_len - max_suffix_len
    if state_budget < 0:
        return None
    clipped = _clip_head_tail(state_ids, state_budget)
    was_truncated = len(clipped) < len(state_ids)
    prefix_ids = bos + state_marker + clipped
    for s in suffixes:
        s["was_truncated"] = was_truncated
        if len(prefix_ids) + len(s["input_ids"]) > max_length:
            return None
    return {"prefix_ids": prefix_ids, "suffixes": suffixes, "was_truncated": was_truncated}