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}