File size: 7,163 Bytes
31f7037
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
"""Validated decision JSONL and the upstream-compatible marker serialization."""
import hashlib
import json
import math
from pathlib import Path

import torch
from torch.utils.data import Dataset

QTYPES = {'choice': 0, 'score': 1, 'noul': 2}


def digest(path):
    with Path(path).open('rb') as stream:
        return hashlib.file_digest(stream, 'sha256').hexdigest()


def validate_row(row, line):
    prefix = f'JSONL line {line}: '
    if not isinstance(row, dict):
        raise ValueError(prefix + 'request must be a JSON object')
    if not isinstance(row.get('state'), (str, dict, list)) or not isinstance(row.get('question'), str):
        raise ValueError(prefix + 'state must be text/JSON and question must be text')
    options = row.get('options')
    if not isinstance(options, list) or not 2 <= len(options) <= 20 or not all(isinstance(x, str) and x for x in options):
        raise ValueError(prefix + 'options must contain 2–20 nonempty rendered descriptions')
    if row.get('type', 'choice') not in QTYPES:
        raise ValueError(prefix + 'type must be choice, score, or noul')
    if row.get('type') == 'noul' and len(options) != 2:
        raise ValueError(prefix + 'noul options must be ordered [false, true]')
    if 'target' in row and (type(row['target']) is not int or not 0 <= row['target'] < len(options)):
        raise ValueError(prefix + 'target must index the supplied option list')
    teacher = row.get('teacher_logits')
    if teacher is not None and (
        not isinstance(teacher, list)
        or len(teacher) != len(options)
        or not all(type(x) in (int, float) and math.isfinite(x) for x in teacher)
    ):
        raise ValueError(prefix + 'teacher logits must be finite and match option count/order')


class Decisions(Dataset):
    def __init__(self, path, require_target=True, require_teacher=False):
        self.rows = []
        with open(path) as stream:
            for line, text in enumerate(stream, 1):
                if not text.strip():
                    continue
                try:
                    row = json.loads(text)
                except json.JSONDecodeError as error:
                    raise ValueError(f'{path}:{line}: invalid JSON: {error.msg}') from error
                validate_row(row, line)
                if require_target and 'target' not in row:
                    raise ValueError(f'{path}:{line}: training requires target')
                if require_teacher and 'teacher_logits' not in row:
                    raise ValueError(f'{path}:{line}: distillation requires teacher_logits')
                self.rows.append(row)
        if not self.rows:
            raise ValueError(f'{path}: dataset contains no rows')
        flags = {'teacher_logits' in row for row in self.rows}
        if len(flags) != 1:
            raise ValueError('Dataset mixes rows with and without teacher logits')

    def __len__(self):
        return len(self.rows)

    def __getitem__(self, index):
        return self.rows[index]


def sequence(tokenizer, row, max_length=8192, head_length=256, *, strict=False):
    if head_length + 4 >= max_length:
        raise ValueError('max_length must leave room beyond the question head')
    if any(x is None for x in (tokenizer.mask_token_id, tokenizer.cls_token_id, tokenizer.sep_token_id)):
        raise ValueError('Tokenizer must define MASK, CLS and SEP IDs')
    state = row['state'] if isinstance(row['state'], str) else json.dumps(row['state'], ensure_ascii=False)
    if strict and any(tokenizer.mask_token in text for text in [state, row['question'], *row['options']]):
        raise ValueError('Reserved model marker in request')
    clean = lambda text: text.replace(tokenizer.mask_token, ' ')
    encode = lambda text: tokenizer(text, add_special_tokens=False)['input_ids']
    head = encode(f"{row.get('type', 'choice')} question: {clean(row['question'])}")
    option_ids = [encode(' ' + clean(x)) for x in row['options']]
    if strict and any(len(x) > 48 for x in option_ids):
        raise ValueError('Option exceeds 48-token model contract')
    options = [[tokenizer.mask_token_id] + x[:48] for x in option_ids]
    budget = head_length - sum(map(len, options))
    if budget < 16:
        per_option = max(4, (head_length - 16) // len(options))
        options = [x[:per_option] for x in options]
        budget = head_length - sum(map(len, options))
    if strict and (len(head) > budget or any(len(x) != len(y) + 1 for x, y in zip(options, option_ids))):
        raise ValueError('Question/options exceed lossless head budget')
    ids = [tokenizer.cls_token_id] + head[:max(8, budget)] + [tokenizer.sep_token_id]
    markers = []
    for option in options:
        markers.append(len(ids))
        ids.extend(option)
    ids.append(tokenizer.sep_token_id)
    state_ids = encode(clean(state))
    room = max_length - len(ids) - 1
    if room < 1:
        raise ValueError('Question/options exceed sequence budget; shorten descriptions')
    if strict and len(state_ids) > room:
        raise ValueError('Game state exceeds lossless context budget')
    result = dict(ids=ids + state_ids[:room] + [tokenizer.sep_token_id], markers=markers,
                  qtype=QTYPES[row.get('type', 'choice')], truncated=len(state_ids) > room)
    if strict:
        result['option_tokens'] = list(map(len, option_ids))
    return result


class Collator:
    def __init__(self, tokenizer, max_length=8192, head_length=256):
        self.tokenizer, self.max_length, self.head_length = tokenizer, max_length, head_length

    def __call__(self, rows, *, include_targets=True):
        encoded = [row['_encoded'] if '_encoded' in row else sequence(
            self.tokenizer, row, self.max_length, self.head_length) for row in rows]
        length = min(self.max_length, ((max(len(x['ids']) for x in encoded) + 7) // 8) * 8)
        count = max(len(x['markers']) for x in encoded)
        ids = torch.full((len(rows), length), self.tokenizer.pad_token_id, dtype=torch.long)
        attention = torch.zeros_like(ids)
        positions = torch.zeros((len(rows), count), dtype=torch.long)
        mask = torch.zeros_like(positions, dtype=torch.bool)
        has_teacher = include_targets and all('teacher_logits' in row for row in rows)
        teacher = torch.zeros_like(positions, dtype=torch.float32) if has_teacher else None
        for i, (row, item) in enumerate(zip(rows, encoded)):
            n, k = len(item['ids']), len(item['markers'])
            ids[i, :n] = torch.tensor(item['ids'])
            attention[i, :n] = 1
            positions[i, :k] = torch.tensor(item['markers'])
            mask[i, :k] = True
            if has_teacher:
                teacher[i, :k] = torch.tensor(row['teacher_logits'])
        batch = dict(input_ids=ids, attention_mask=attention, marker_pos=positions,
                     marker_mask=mask, qtype=torch.tensor([x['qtype'] for x in encoded]))
        if include_targets and all('target' in row for row in rows):
            batch['labels'] = torch.tensor([row['target'] for row in rows])
        if has_teacher:
            batch['teacher_logits'] = teacher
        return batch