File size: 3,812 Bytes
c800ccd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
"""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