File size: 6,340 Bytes
26b3403
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
144
145
146
147
148
149
150
151
152
"""Encoder rendering and windowing shared by training, evaluation and the Core ML engine.

A window is one encoder sequence:

    [CLS] <type> <question> [SEP] ([MASK] <option>)* [SEP] <context slice> [SEP]

Each option is read at its [MASK] marker. Nothing is truncated: long questions keep a
head and tail in every window and their middle joins the scanned context; option lists
that do not fit are packed into several groups; an option too long for one window is
split into pieces, each with its own marker; the context (question middle + state) is
scanned in overlapping slices. An option's logit is the log-mean-exp over every marker
occurrence of it (all windows and pieces), so training and inference pool identically.
"""

from __future__ import annotations

import json
from dataclasses import dataclass, field
from typing import Any

from tokenizers import Tokenizer

CLS, SEP, PAD, MASK = 50281, 50282, 50283, 50284
MAX_OPTIONS = 255
TYPE_TEXT = {"choice": "choice:", "noul": "yes or no:", "score": "rate on the scale:"}


def describe(value: Any) -> str:
    return value if isinstance(value, str) else json.dumps(value, ensure_ascii=False)


def option_list(question: dict[str, Any]) -> tuple[list[str], list[str]]:
    """Same keys/descriptions as the server's reference `option_list`."""
    kind = question["type"]
    criteria = question.get("criteria")
    if kind == "choice":
        keys = list(criteria)
        return keys, [k if v is None else f"{k}: {describe(v)}" for k, v in criteria.items()]
    if kind == "noul":
        c = criteria or {}
        return ["false", "true"], [describe(c.get("false") or "No / false"), describe(c.get("true") or "Yes / true")]
    if kind == "score":
        return [str(i) for i in range(len(criteria))], [describe(v) for v in criteria]
    raise ValueError(f"Unknown question type: {kind}")


@dataclass(frozen=True)
class WindowConfig:
    max_len: int = 512
    buckets: tuple[int, ...] = (128, 256, 512)
    max_slots: int = 64  # option markers per window
    q_head: int = 64
    q_tail: int = 128
    min_context: int = 96  # context tokens reserved per window when options are packed
    overlap: float = 0.25


@dataclass
class Window:
    ids: list[int]
    positions: list[int] = field(default_factory=list)  # marker token positions
    options: list[int] = field(default_factory=list)  # option index per marker

    @property
    def length(self) -> int:
        return len(self.ids)


class Unsupported(ValueError):
    pass


class Renderer:
    def __init__(self, tokenizer_path: str, config: WindowConfig = WindowConfig()):
        self.tok = Tokenizer.from_file(tokenizer_path)
        self.tok.no_padding()
        self.tok.no_truncation()
        self.cfg = config
        self._type_ids = {k: self._enc(v) for k, v in TYPE_TEXT.items()}

    def _enc(self, text: str) -> list[int]:
        return self.tok.encode(text, add_special_tokens=False).ids if text else []

    def _enc_batch(self, texts: list[str]) -> list[list[int]]:
        return [e.ids for e in self.tok.encode_batch(texts, add_special_tokens=False)] if texts else []

    def bucket(self, n: int) -> int:
        for b in self.cfg.buckets:
            if n <= b:
                return b
        raise Unsupported(f"window of {n} tokens exceeds {self.cfg.buckets[-1]}")

    def windows(self, state: Any, question: dict[str, Any]) -> tuple[list[Window], int]:
        """All windows for one question and its option count."""
        cfg = self.cfg
        _, descriptions = option_list(question)
        n_opt = len(descriptions)
        if not 1 <= n_opt <= MAX_OPTIONS:
            raise Unsupported(f"{n_opt} options exceeds the declared limit of {MAX_OPTIONS}")
        state_text = describe(state) if state not in (None, "", {}) else ""
        instr = describe(question.get("instructions") or "Choose the best matching option.")
        q_ids, s_ids, *o_ids = self._enc_batch([instr, state_text] + descriptions)
        # question: keep head + tail, move the middle into the scanned context
        if len(q_ids) > cfg.q_head + cfg.q_tail:
            middle = q_ids[cfg.q_head:len(q_ids) - cfg.q_tail]
            q_ids = q_ids[:cfg.q_head] + q_ids[len(q_ids) - cfg.q_tail:]
            context = middle + (s_ids if not s_ids else [SEP] + s_ids)
        else:
            context = s_ids
        prefix = [CLS] + self._type_ids[question["type"]] + q_ids + [SEP]
        fixed = len(prefix) + 2  # SEP after options, SEP at end
        opt_budget = cfg.max_len - fixed - (cfg.min_context if context else 0)
        if opt_budget < 8:
            raise Unsupported("question head/tail leaves no room for options")
        # option pieces: (option index, tokens), each piece fits opt_budget with its marker
        pieces: list[tuple[int, list[int]]] = []
        for i, ids in enumerate(o_ids):
            ids = ids or [PAD]  # empty description still gets a marker
            step = opt_budget - 1
            for s in range(0, len(ids), step):
                pieces.append((i, ids[s:s + step]))
        groups: list[list[tuple[int, list[int]]]] = [[]]
        used = 0
        for piece in pieces:
            cost = 1 + len(piece[1])
            if groups[-1] and (used + cost > opt_budget or len(groups[-1]) >= cfg.max_slots):
                groups.append([])
                used = 0
            groups[-1].append(piece)
            used += cost
        out: list[Window] = []
        for group in groups:
            body = list(prefix)
            positions, options = [], []
            for i, ids in group:
                positions.append(len(body))
                options.append(i)
                body.append(MASK)
                body.extend(ids)
            body.append(SEP)
            room = cfg.max_len - len(body) - 1
            if not context:
                out.append(Window(body + [SEP], positions, options))
                continue
            if room <= 0:
                raise Unsupported("no room for context")
            stride = max(1, int(room * (1 - cfg.overlap)))
            starts = [0] if len(context) <= room else list(range(0, len(context) - room, stride)) + [len(context) - room]
            for s in starts:
                out.append(Window(body + context[s:s + room] + [SEP], list(positions), list(options)))
        return out, n_opt