openjev-e4b / eval /harness /src /formatting.py
bambamdevs's picture
Publish OpenJEV E4B 1.0
03223d7
Raw History Blame Contribute Delete
8.83 kB
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}