Download data_utils.py from HrushikeshGangane/decision_maker: direct link, hf CLI and curl.
- Browser
- Download file 3.81 kB
-
https://huggingface.co/HrushikeshGangane/decision_maker/resolve/main/data_utils.py
- Command line
-
hf download hf://HrushikeshGangane/decision_maker/data_utils.py
-
curl -L -o data_utils.py https://huggingface.co/HrushikeshGangane/decision_maker/resolve/main/data_utils.py
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 | |