decision_maker / data_utils.py
HrushikeshGangane's picture
Publish decision_maker inference bundle
c800ccd verified
Raw History Blame Contribute Delete
3.81 kB
"""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', '<unknown>')} 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