"""JSONL dataset, marker-aware encoding, and batch collation.""" from __future__ import annotations import json from pathlib import Path from typing import Any import torch from torch.utils.data import Dataset from decision_model import SPECIAL_TOKENS class DecisionDataset(Dataset): def __init__(self, path: str | Path, tokenizer=None, max_length: int | None = None): with open(path, encoding="utf-8") as f: self.rows = [json.loads(line) for line in f if line.strip()] for row in self.rows: if not row.get("options") or not 0 <= row.get("label", -1) < len(row["options"]): raise ValueError(f"Bad decision example: {row.get('id', row)}") self.total_rows = len(self.rows) self.skipped_oversize = 0 if tokenizer is not None and max_length is not None: kept = [] for row in self.rows: if len(encode_row(tokenizer, row)["input_ids"]) <= max_length: kept.append(row) else: self.skipped_oversize += 1 self.rows = kept if not self.rows: raise ValueError(f"Every example in {path} exceeds max_length={max_length}") def __len__(self): return len(self.rows) def __getitem__(self, index: int) -> dict[str, Any]: return self.rows[index] def encode_row(tokenizer, row: dict[str, Any]) -> dict[str, Any]: """Tokenize one row without padding, retaining structural marker positions.""" marker_ids = {token: tokenizer.convert_tokens_to_ids(token) for token in SPECIAL_TOKENS} ids = [tokenizer.cls_token_id] ids += [marker_ids["[STATE]"]] + tokenizer.encode(row.get("state", ""), add_special_tokens=False) ids += [marker_ids["[QUESTION]"]] + tokenizer.encode(row["question"], add_special_tokens=False) options = [] for option in row["options"]: options.append(len(ids)) ids += [marker_ids["[OPTION]"]] + tokenizer.encode(str(option), add_special_tokens=False) decide = len(ids) ids += [marker_ids["[DECIDE]"], tokenizer.sep_token_id] return {"input_ids": ids, "decide_position": decide, "option_positions": options} def make_collate(tokenizer, max_length: int): def collate(rows: list[dict[str, Any]]) -> dict[str, Any]: # Tokenize fields separately so marker locations stay exact even under truncation. sequences, decide_positions, option_positions = [], [], [] for row in rows: encoded = encode_row(tokenizer, row) ids, decide, opts = encoded["input_ids"], encoded["decide_position"], encoded["option_positions"] if len(ids) > max_length: raise ValueError(f"Oversize example reached collate: {row.get('id', '')} exceeds max_length={max_length}. Construct DecisionDataset with tokenizer/max_length to filter it.") sequences.append(ids) decide_positions.append(decide) option_positions.append(opts) padded = tokenizer.pad({"input_ids": sequences}, padding=True, return_tensors="pt") max_options = max(len(p) for p in option_positions) pos = torch.full((len(rows), max_options), -1, dtype=torch.long) mask = torch.zeros((len(rows), max_options), dtype=torch.bool) for i, positions in enumerate(option_positions): pos[i, : len(positions)] = torch.tensor(positions) mask[i, : len(positions)] = True return { **padded, "decide_positions": torch.tensor(decide_positions, dtype=torch.long), "option_positions": pos, "option_mask": mask, "labels": torch.tensor([row["label"] for row in rows], dtype=torch.long), "rows": rows, } return collate