sharada 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
sharada/__init__.py ADDED
@@ -0,0 +1,31 @@
1
+ """Sharada — typed decisions from text in one forward pass.
2
+
3
+ from sharada import DecisionModel
4
+
5
+ model = DecisionModel.from_pretrained("lenabarretta/sharada-base")
6
+ d = model.decide("My card hasn't arrived yet",
7
+ "Which team should handle this?",
8
+ ["billing", "technical", "sales"])
9
+ d.answer, d.confidence
10
+
11
+ The model reads the text, the question and every option in one pass and returns a probability for each
12
+ option. Every option is a branch of the sequence starting at the same position, so their order cannot
13
+ change the answer; an option reads the text, the question and itself, so its score does not depend on
14
+ which other options are offered.
15
+ """
16
+
17
+ from .calibrate import calibrate, check_passport, fit_temperature, save_passport
18
+ from .data import Example, Request
19
+ from .evaluate import calibration_error, evaluate, latency, risk_coverage
20
+ from .layout import Limits
21
+ from .model import Config, Decision, DecisionModel
22
+ from .policy import Action, Policy, escalation_budget
23
+ from .train import fit
24
+
25
+ __all__ = [
26
+ "DecisionModel", "Config", "Decision", "Request", "Example", "Limits",
27
+ "fit", "calibrate", "fit_temperature", "save_passport", "check_passport",
28
+ "evaluate", "calibration_error", "risk_coverage", "latency",
29
+ "Policy", "Action", "escalation_budget",
30
+ ]
31
+ __version__ = "0.1.0"
sharada/calibrate.py ADDED
@@ -0,0 +1,115 @@
1
+ """Calibration, and the passport that says where it holds.
2
+
3
+ A model is not calibrated in general; it is calibrated on a distribution, and it goes stale when the
4
+ traffic moves. So fitting a temperature also writes a passport: what it was fitted on, how far off the
5
+ probabilities were, and what should make you fit it again.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import datetime as dt
11
+ import json
12
+ import math
13
+ import pathlib
14
+
15
+ from .data import Example
16
+
17
+ VALID_DAYS = 90
18
+
19
+
20
+ def _negative_log_likelihood(log_probabilities: list[list[float]], labels: list[int], temperature: float) -> float:
21
+ total = 0.0
22
+ for row, label in zip(log_probabilities, labels):
23
+ scaled = [v / temperature for v in row]
24
+ top = max(scaled)
25
+ total += -(scaled[label] - (top + math.log(sum(math.exp(v - top) for v in scaled))))
26
+ return total / len(labels)
27
+
28
+
29
+ def fit_temperature(probabilities: list[list[float]], labels: list[int],
30
+ grid: tuple[float, float, int] = (0.25, 8.0, 61)) -> float:
31
+ """One number per task: divide the scores by it. Chosen to minimise the log loss on held-out
32
+ answers, which is what a temperature can honestly be chosen on."""
33
+ low, high, steps = grid
34
+ log_probabilities = [[math.log(max(p, 1e-12)) for p in row] for row in probabilities]
35
+ candidates = [math.exp(math.log(low) + (math.log(high) - math.log(low)) * i / (steps - 1)) for i in range(steps)]
36
+ return min(candidates, key=lambda t: _negative_log_likelihood(log_probabilities, labels, t))
37
+
38
+
39
+ def fingerprint(model, examples: list[Example]) -> dict:
40
+ """What the examples looked like, so a later drift is visible."""
41
+ out: dict[str, dict] = {}
42
+ for task in sorted({e.task for e in examples}):
43
+ subset = [e for e in examples if e.task == task]
44
+ lengths = sorted(len(model.layout(e.text, e.question, e.options).ids) for e in subset)
45
+ counts: dict[str, int] = {}
46
+ for e in subset:
47
+ counts[e.options[e.label]] = counts.get(e.options[e.label], 0) + 1
48
+ out[task] = {"n": len(subset), "options": subset[0].options,
49
+ "median_tokens": lengths[len(lengths) // 2],
50
+ "answers": counts}
51
+ return out
52
+
53
+
54
+ def calibrate(model, examples: list[Example], batch_size: int = 32, valid_days: int = VALID_DAYS) -> dict:
55
+ """Fit one temperature per task on these examples and return the passport. The temperatures are
56
+ stored on the model, so `decide` uses them from then on."""
57
+ from .evaluate import evaluate
58
+
59
+ raw = model.probabilities(examples, batch_size=batch_size, calibrated=False)
60
+ for task in sorted({e.task for e in examples}):
61
+ index = [i for i, e in enumerate(examples) if e.task == task]
62
+ model.config.temperatures[task] = fit_temperature([raw[i] for i in index],
63
+ [examples[i].label for i in index])
64
+ measured = evaluate(model, examples, batch_size=batch_size, calibrated=True)
65
+ today = dt.date.today()
66
+ return {
67
+ "created": today.isoformat(),
68
+ "valid_until": (today + dt.timedelta(days=valid_days)).isoformat(),
69
+ "encoder": model.config.encoder,
70
+ "examples": len(examples),
71
+ "temperatures": dict(model.config.temperatures),
72
+ "fingerprint": fingerprint(model, examples),
73
+ "calibration_error": measured["calibration_error"],
74
+ "calibration_error_interval": measured["calibration_error_interval"],
75
+ "accuracy": measured["accuracy"],
76
+ "reliability": measured["reliability"],
77
+ "risk_coverage": measured["risk_coverage"],
78
+ "recalibrate_if": {
79
+ "after": "valid_until",
80
+ "options_change": True,
81
+ "median_tokens_shift": 0.5, # half as long, or twice as long
82
+ },
83
+ }
84
+
85
+
86
+ def save_passport(passport: dict, path: str | pathlib.Path) -> pathlib.Path:
87
+ path = pathlib.Path(path)
88
+ path.write_text(json.dumps(passport, indent=1))
89
+ return path
90
+
91
+
92
+ def check_passport(passport: dict, examples: list[Example] | None = None, model=None,
93
+ today: dt.date | None = None) -> list[str]:
94
+ """What is wrong with trusting these probabilities today. An empty list means nothing found —
95
+ run it in the deployment pipeline and fail the build on anything it returns."""
96
+ today = today or dt.date.today()
97
+ problems = []
98
+ if today.isoformat() > passport["valid_until"]:
99
+ problems.append(f"the calibration expired on {passport['valid_until']}; fit it again on recent answers")
100
+ if passport["examples"] < 200:
101
+ problems.append(f"fitted on {passport['examples']} answers — the interval on the calibration error is wide")
102
+ if examples and model:
103
+ now = fingerprint(model, examples)
104
+ for task, then in passport["fingerprint"].items():
105
+ if task not in now:
106
+ continue
107
+ if now[task]["options"] != then["options"]:
108
+ problems.append(f"{task}: the options have changed since the calibration")
109
+ ratio = now[task]["median_tokens"] / max(then["median_tokens"], 1)
110
+ if ratio < 0.5 or ratio > 2:
111
+ problems.append(f"{task}: texts are now {ratio:.1f}× the length they were calibrated on")
112
+ for task in now:
113
+ if task not in passport["fingerprint"]:
114
+ problems.append(f"{task}: never calibrated — its probabilities are raw")
115
+ return problems
sharada/data.py ADDED
@@ -0,0 +1,105 @@
1
+ """Requests, examples and batches."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import random
6
+ from dataclasses import dataclass, field
7
+
8
+ import torch
9
+
10
+ from .layout import FIRST_OPTION, Laid, Layout
11
+
12
+ KINDS = ("choice", "scale", "binary") # unordered labels, an ordered scale, yes or no
13
+
14
+
15
+ @dataclass
16
+ class Request:
17
+ """One question about one text."""
18
+
19
+ text: str
20
+ question: str
21
+ options: list[str]
22
+ kind: str = "choice"
23
+ task: str = "task" # a name, so temperatures can be fitted per task
24
+
25
+ def __post_init__(self):
26
+ if self.kind not in KINDS:
27
+ raise ValueError(f"kind must be one of {KINDS}, got {self.kind!r}")
28
+
29
+
30
+ @dataclass
31
+ class Example(Request):
32
+ """A request whose answer is known: `label` is the index of the right option."""
33
+
34
+ label: int = 0
35
+
36
+ def __post_init__(self):
37
+ super().__post_init__()
38
+ if not 0 <= self.label < len(self.options):
39
+ raise ValueError(f"label {self.label} is outside the {len(self.options)} options")
40
+
41
+
42
+ @dataclass
43
+ class Batch:
44
+ ids: torch.Tensor
45
+ positions: torch.Tensor
46
+ segments: torch.Tensor
47
+ present: torch.Tensor # [batch, options] — which option slots are real
48
+ kinds: torch.Tensor
49
+ n_options: torch.Tensor
50
+ labels: torch.Tensor | None = None
51
+ tasks: list[str] = field(default_factory=list)
52
+ width: int = 0 # the widest number of options in this batch
53
+
54
+ def to(self, device) -> "Batch":
55
+ move = lambda v: v.to(device) if torch.is_tensor(v) else v
56
+ return Batch(move(self.ids), move(self.positions), move(self.segments), move(self.present),
57
+ move(self.kinds), move(self.n_options), move(self.labels), self.tasks, self.width)
58
+
59
+
60
+ def collate(laid: list[Laid], requests: list[Request], pad_id: int) -> Batch:
61
+ length = max(len(l) for l in laid)
62
+ width = max(l.n_options for l in laid)
63
+ n = len(laid)
64
+ ids = torch.full((n, length), pad_id, dtype=torch.long)
65
+ positions = torch.zeros((n, length), dtype=torch.long)
66
+ segments = torch.full((n, length), -1, dtype=torch.long)
67
+ for i, l in enumerate(laid):
68
+ k = len(l)
69
+ ids[i, :k] = torch.tensor(l.ids)
70
+ positions[i, :k] = torch.tensor(l.positions)
71
+ segments[i, :k] = torch.tensor(l.segments)
72
+ counts = torch.tensor([l.n_options for l in laid])
73
+ labels = None
74
+ if all(isinstance(r, Example) for r in requests):
75
+ labels = torch.tensor([r.label for r in requests])
76
+ return Batch(ids=ids, positions=positions, segments=segments,
77
+ present=torch.arange(width)[None] < counts[:, None],
78
+ kinds=torch.tensor([KINDS.index(r.kind) for r in requests]),
79
+ n_options=counts, labels=labels, tasks=[r.task for r in requests], width=width)
80
+
81
+
82
+ def batches(requests: list[Request], layout: Layout, batch_size: int, pad_id: int,
83
+ shuffle: bool = False, seed: int = 0):
84
+ """Mini-batches. When shuffling, requests of similar length travel together, so a batch of short
85
+ questions is not padded out to the length of one with sixty options."""
86
+ order = list(range(len(requests)))
87
+ laid = [layout(r.text, r.question, r.options) for r in requests]
88
+ if shuffle:
89
+ rng = random.Random(seed)
90
+ rng.shuffle(order)
91
+ window = batch_size * 50
92
+ order = [j for w in range(0, len(order), window)
93
+ for j in sorted(order[w:w + window], key=lambda j: len(laid[j]))]
94
+ chunks = [order[i:i + batch_size] for i in range(0, len(order), batch_size)]
95
+ if shuffle:
96
+ random.Random(seed + 1).shuffle(chunks)
97
+ for chunk in chunks:
98
+ yield chunk, collate([laid[j] for j in chunk], [requests[j] for j in chunk], pad_id)
99
+
100
+
101
+ def option_mass(segments: torch.Tensor, states: torch.Tensor, width: int) -> torch.Tensor:
102
+ """Mean of each option's own final token states -> [batch, options, hidden]."""
103
+ index = FIRST_OPTION + torch.arange(width, device=states.device)
104
+ member = (segments[:, :, None] == index[None, None, :]).to(states.dtype)
105
+ return torch.einsum("blk,bld->bkd", member, states) / member.sum(1).clamp(min=1)[..., None]
sharada/evaluate.py ADDED
@@ -0,0 +1,105 @@
1
+ """Measuring a decision model: the answer, the stated number, and the whole distribution."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import math
6
+ import random
7
+ import time
8
+
9
+ from .data import Example
10
+
11
+
12
+ def _bands(confidence: list[float], n: int) -> list[int]:
13
+ return [min(int(c * n), n - 1) for c in confidence]
14
+
15
+
16
+ def calibration_error(confidence: list[float], correct: list[int], bands: int = 15) -> float:
17
+ """The average gap between the probability stated and how often it turns out right, over equal
18
+ bands, weighted by how many answers fall in each. Zero means this estimate finds no gap — on a
19
+ finite sample that is not proof of none."""
20
+ index = _bands(confidence, bands)
21
+ total = 0.0
22
+ for b in range(bands):
23
+ inside = [i for i, j in enumerate(index) if j == b]
24
+ if inside:
25
+ said = sum(confidence[i] for i in inside) / len(inside)
26
+ was = sum(correct[i] for i in inside) / len(inside)
27
+ total += len(inside) / len(confidence) * abs(said - was)
28
+ return total
29
+
30
+
31
+ def calibration_interval(confidence: list[float], correct: list[int], bands: int = 15,
32
+ draws: int = 200, seed: int = 0) -> tuple[float, float]:
33
+ """A bootstrap interval for the calibration error — on a few hundred answers it is wide, and
34
+ saying so is the point."""
35
+ rng = random.Random(seed)
36
+ n = len(confidence)
37
+ values = []
38
+ for _ in range(draws):
39
+ pick = [rng.randrange(n) for _ in range(n)]
40
+ values.append(calibration_error([confidence[i] for i in pick], [correct[i] for i in pick], bands))
41
+ values.sort()
42
+ return values[int(0.025 * draws)], values[int(0.975 * draws) - 1]
43
+
44
+
45
+ def reliability_curve(confidence: list[float], correct: list[int], bands: int = 10) -> list[dict]:
46
+ index = _bands(confidence, bands)
47
+ curve = []
48
+ for b in range(bands):
49
+ inside = [i for i, j in enumerate(index) if j == b]
50
+ curve.append({"from": b / bands, "n": len(inside),
51
+ "stated": sum(confidence[i] for i in inside) / len(inside) if inside else None,
52
+ "observed": sum(correct[i] for i in inside) / len(inside) if inside else None})
53
+ return curve
54
+
55
+
56
+ def risk_coverage(confidence: list[float], correct: list[int], steps: int = 20) -> list[dict]:
57
+ """Answer only where the model is surest: what accuracy, at what share of the traffic."""
58
+ order = sorted(range(len(confidence)), key=lambda i: -confidence[i])
59
+ out = []
60
+ for k in range(1, steps + 1):
61
+ take = max(1, round(len(order) * k / steps))
62
+ kept = [correct[i] for i in order[:take]]
63
+ out.append({"coverage": take / len(order), "accuracy": sum(kept) / len(kept),
64
+ "threshold": confidence[order[take - 1]]})
65
+ return out
66
+
67
+
68
+ def evaluate(model, examples: list[Example], batch_size: int = 32, calibrated: bool = True) -> dict:
69
+ """Everything above, on one set of labelled examples."""
70
+ probabilities = model.probabilities(examples, batch_size=batch_size, calibrated=calibrated)
71
+ confidence, correct, true_p, squared = [], [], [], []
72
+ for p, e in zip(probabilities, examples):
73
+ best = max(range(len(p)), key=p.__getitem__)
74
+ confidence.append(p[best])
75
+ correct.append(int(best == e.label))
76
+ true_p.append(p[e.label])
77
+ squared.append(sum((q - (1.0 if i == e.label else 0.0)) ** 2 for i, q in enumerate(p)))
78
+ low, high = calibration_interval(confidence, correct)
79
+ return {"n": len(examples),
80
+ "accuracy": sum(correct) / len(correct),
81
+ "mean_confidence": sum(confidence) / len(confidence),
82
+ "calibration_error": calibration_error(confidence, correct),
83
+ "calibration_error_interval": [low, high],
84
+ "log_loss": sum(-math.log(max(p, 1e-9)) for p in true_p) / len(true_p),
85
+ "brier": sum(squared) / len(squared),
86
+ "reliability": reliability_curve(confidence, correct),
87
+ "risk_coverage": risk_coverage(confidence, correct)}
88
+
89
+
90
+ def latency(model, request, batch_size: int = 1, repeats: int = 20) -> dict:
91
+ """Milliseconds per call and per question, as the model is loaded right now."""
92
+ import torch
93
+
94
+ requests = [request] * batch_size
95
+ for _ in range(3):
96
+ model.probabilities(requests, batch_size=batch_size)
97
+ if next(model.parameters()).is_cuda:
98
+ torch.cuda.synchronize()
99
+ start = time.time()
100
+ for _ in range(repeats):
101
+ model.probabilities(requests, batch_size=batch_size)
102
+ if next(model.parameters()).is_cuda:
103
+ torch.cuda.synchronize()
104
+ per_call = (time.time() - start) / repeats * 1e3
105
+ return {"batch": batch_size, "ms_per_call": per_call, "ms_per_question": per_call / batch_size}
sharada/layout.py ADDED
@@ -0,0 +1,78 @@
1
+ """One request, laid out as one sequence.
2
+
3
+ position: 0 1 … n n+1 n+2 … m m+1 │ m+2 … │ m+2 …
4
+ token: [CLS] the text [SEP] question [SEP]│ option 0 [SEP] │ option 1 [SEP]
5
+
6
+ Every option is a branch of its own and every branch starts at the same position, the one right after
7
+ the question. The encoder's positions are rotary, so to it no option is earlier or later than another:
8
+ the order of the options cannot change their scores. `masking.who_reads_whom` then decides who may read
9
+ whom; together the two give the properties the model is built on.
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ import functools
15
+ from dataclasses import dataclass, field
16
+
17
+ TEXT = 0
18
+ QUESTION = 1
19
+ FIRST_OPTION = 2 # option i has segment FIRST_OPTION + i
20
+
21
+
22
+ @dataclass(frozen=True)
23
+ class Limits:
24
+ """How many tokens each part of a request may use."""
25
+
26
+ text: int = 256
27
+ question: int = 48
28
+ option: int = 12
29
+
30
+ def as_dict(self) -> dict:
31
+ return {"text": self.text, "question": self.question, "option": self.option}
32
+
33
+
34
+ @dataclass
35
+ class Laid:
36
+ """A request turned into token ids, position ids and one segment per token."""
37
+
38
+ ids: list[int]
39
+ positions: list[int]
40
+ segments: list[int]
41
+ n_options: int
42
+ option_start: int = field(default=0)
43
+
44
+ def __len__(self) -> int:
45
+ return len(self.ids)
46
+
47
+
48
+ class Layout:
49
+ """Lays requests out for one tokenizer. Token ids of repeated strings — option names, instructions,
50
+ the same text asked about twice — are cached, which is most of the cost of preparing a batch."""
51
+
52
+ def __init__(self, tokenizer, limits: Limits | None = None):
53
+ self.tokenizer = tokenizer
54
+ self.limits = limits or Limits()
55
+ self._ids = functools.lru_cache(maxsize=100_000)(self._encode)
56
+
57
+ def _encode(self, text: str) -> tuple[int, ...]:
58
+ return tuple(self.tokenizer(text, add_special_tokens=False)["input_ids"])
59
+
60
+ def __call__(self, text: str, question: str, options: list[str]) -> Laid:
61
+ if len(options) < 2:
62
+ raise ValueError("a decision needs at least two options")
63
+ cls_id, sep_id = self.tokenizer.cls_token_id, self.tokenizer.sep_token_id
64
+ body = list(self._ids(text))[: self.limits.text]
65
+ ask = list(self._ids(question))[: self.limits.question]
66
+
67
+ ids = [cls_id] + body + [sep_id] + ask + [sep_id]
68
+ segments = [TEXT] * (len(body) + 2) + [QUESTION] * (len(ask) + 1)
69
+ positions = list(range(len(ids)))
70
+
71
+ start = len(ids)
72
+ for i, option in enumerate(options):
73
+ tokens = list(self._ids(" " + option))[: self.limits.option] + [sep_id]
74
+ ids += tokens
75
+ segments += [FIRST_OPTION + i] * len(tokens)
76
+ positions += range(start, start + len(tokens)) # every option restarts here
77
+ return Laid(ids=ids, positions=positions, segments=segments, n_options=len(options),
78
+ option_start=start)
sharada/masking.py ADDED
@@ -0,0 +1,48 @@
1
+ """Who reads whom.
2
+
3
+ The text reads only the text, so its states do not depend on the question or the options — a text asked
4
+ several questions could be read once. The question reads the text and itself. An option reads the text,
5
+ the question and its own tokens, and nothing else: with seventy-seven options on offer, each one still
6
+ sees only the message and itself, and its score does not depend on which other options came along.
7
+
8
+ `read_each_other=True` keeps the other behaviour — options that also read one another — which is what
9
+ this model was first trained with. On label sets of a few options it costs nothing; on seventy-seven it
10
+ cost thirty points of accuracy, which is why it is not the default.
11
+
12
+ The encoder's local-attention layers only look a fixed distance away, and the distance is measured in
13
+ positions, so an option sees the end of the question the same way wherever it sits in the sequence.
14
+ """
15
+
16
+ from __future__ import annotations
17
+
18
+ import torch
19
+
20
+ from .layout import QUESTION, TEXT
21
+
22
+
23
+ def who_reads_whom(segments: torch.Tensor, positions: torch.Tensor, local_reach: int,
24
+ read_each_other: bool = False) -> tuple[torch.Tensor, torch.Tensor]:
25
+ """-> two [batch, length, length] boolean masks: for the global layers and for the local ones.
26
+ True at (i, j) means query token i may read key token j."""
27
+ q, k = segments[:, :, None], segments[:, None, :]
28
+ real = k >= 0 # padding is not a key
29
+ context = (k == TEXT) | (k == QUESTION)
30
+ options_see = real if read_each_other else (context | (k == q))
31
+ full = torch.where(q == TEXT, k == TEXT, torch.where(q == QUESTION, context, options_see)) & real
32
+ full = full | torch.eye(segments.size(1), dtype=torch.bool, device=segments.device)[None]
33
+ near = (positions[:, :, None] - positions[:, None, :]).abs() <= local_reach
34
+ return full, full & near
35
+
36
+
37
+ def encoder_masks(segments: torch.Tensor, positions: torch.Tensor, local_reach: int,
38
+ attn_implementation: str, dtype: torch.dtype, read_each_other: bool = False) -> dict:
39
+ """The masks in the form the encoder expects: booleans for fused attention, an additive float mask
40
+ for the plain implementation."""
41
+ out = {}
42
+ names = ("full_attention", "sliding_attention")
43
+ for name, allowed in zip(names, who_reads_whom(segments, positions, local_reach, read_each_other)):
44
+ mask = allowed[:, None]
45
+ if attn_implementation == "eager":
46
+ mask = torch.zeros(mask.shape, dtype=dtype, device=mask.device).masked_fill(~mask, torch.finfo(dtype).min)
47
+ out[name] = mask
48
+ return out
sharada/model.py ADDED
@@ -0,0 +1,167 @@
1
+ """The decision model: an encoder, a small read-out, one number per option."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import pathlib
7
+ from dataclasses import asdict, dataclass, field
8
+
9
+ import torch
10
+ import torch.nn as nn
11
+
12
+ from .data import Batch, Example, Request, batches, collate, option_mass
13
+ from .layout import Laid, Layout, Limits
14
+
15
+ DEFAULT_ENCODER = "answerdotai/ModernBERT-base"
16
+
17
+
18
+ @dataclass
19
+ class Config:
20
+ encoder: str = DEFAULT_ENCODER
21
+ limits: dict = field(default_factory=lambda: Limits().as_dict())
22
+ local_reach: int | None = None # taken from the encoder when left out
23
+ read_each_other: bool = False # options reading one another; off by default, see masking.py
24
+ temperatures: dict = field(default_factory=dict) # per task, fitted by `calibrate`
25
+ version: int = 1
26
+
27
+
28
+ @dataclass
29
+ class Decision:
30
+ """What the model answers, and how sure it says it is."""
31
+
32
+ answer: str
33
+ index: int
34
+ probabilities: dict[str, float]
35
+ confidence: float
36
+
37
+ def __repr__(self) -> str:
38
+ return f"Decision({self.answer!r}, confidence={self.confidence:.3f})"
39
+
40
+
41
+ class DecisionModel(nn.Module):
42
+ def __init__(self, config: Config | None = None, tokenizer=None, encoder=None):
43
+ super().__init__()
44
+ from transformers import AutoModel, AutoTokenizer
45
+
46
+ self.config = config or Config()
47
+ self.tokenizer = tokenizer or AutoTokenizer.from_pretrained(self.config.encoder)
48
+ self.attn = "eager" if not torch.cuda.is_available() else "sdpa"
49
+ self.encoder = encoder or AutoModel.from_pretrained(self.config.encoder, attn_implementation=self.attn)
50
+ hidden = self.encoder.config.hidden_size
51
+ self.read = nn.Sequential(nn.Linear(2 * hidden, hidden), nn.GELU(), nn.LayerNorm(hidden),
52
+ nn.Linear(hidden, 1))
53
+ if self.config.local_reach is None:
54
+ self.config.local_reach = getattr(self.encoder.config, "local_attention", 128) // 2
55
+ self.layout = Layout(self.tokenizer, Limits(**self.config.limits))
56
+
57
+ # ── the forward pass ──────────────────────────────────────────────────────────────────────
58
+ def forward(self, batch: Batch) -> torch.Tensor:
59
+ """-> [batch, options] scores; options that do not exist cannot win."""
60
+ from .masking import encoder_masks
61
+
62
+ dtype = self.encoder.embeddings.tok_embeddings.weight.dtype
63
+ masks = encoder_masks(batch.segments, batch.positions, self.config.local_reach,
64
+ self.attn, dtype, self.config.read_each_other)
65
+ states = self.encoder(input_ids=batch.ids, position_ids=batch.positions,
66
+ attention_mask=masks).last_hidden_state
67
+ options = option_mass(batch.segments, states, batch.width)
68
+ summary = states[:, :1].expand(-1, batch.width, -1)
69
+ scores = self.read(torch.cat([options, options * summary], -1)).squeeze(-1).float()
70
+ return scores.masked_fill(~batch.present, -1e4)
71
+
72
+ # ── asking it things ──────────────────────────────────────────────────────────────────────
73
+ @torch.no_grad()
74
+ def probabilities(self, requests: list[Request], batch_size: int = 32,
75
+ calibrated: bool = True) -> list[list[float]]:
76
+ """One distribution over the options of each request, in the order they were given."""
77
+ self.eval()
78
+ device = next(self.parameters()).device
79
+ out: list[list[float]] = [[] for _ in requests]
80
+ for chunk, batch in batches(requests, self.layout, batch_size, self.tokenizer.pad_token_id):
81
+ scores = self(batch.to(device))
82
+ for row, j in enumerate(chunk):
83
+ k = int(batch.n_options[row])
84
+ z = scores[row, :k].float().cpu()
85
+ if calibrated:
86
+ z = z / self.temperature_for(requests[j])
87
+ out[j] = torch.softmax(z, -1).tolist()
88
+ return out
89
+
90
+ def decide(self, text: str, question: str, options: list[str], kind: str = "choice",
91
+ task: str = "task", calibrated: bool = True) -> Decision:
92
+ """One decision, start to finish."""
93
+ request = Request(text=text, question=question, options=options, kind=kind, task=task)
94
+ return self.decide_many([request], calibrated=calibrated)[0]
95
+
96
+ def decide_many(self, requests: list[Request], batch_size: int = 32,
97
+ calibrated: bool = True) -> list[Decision]:
98
+ decisions = []
99
+ for request, probabilities in zip(requests, self.probabilities(requests, batch_size, calibrated)):
100
+ best = max(range(len(probabilities)), key=probabilities.__getitem__)
101
+ decisions.append(Decision(answer=request.options[best], index=best,
102
+ probabilities=dict(zip(request.options, probabilities)),
103
+ confidence=probabilities[best]))
104
+ return decisions
105
+
106
+ def temperature_for(self, request: Request) -> float:
107
+ return float(self.config.temperatures.get(request.task, 1.0))
108
+
109
+ # ── saving and loading ────────────────────────────────────────────────────────────────────
110
+ def save(self, path: str | pathlib.Path) -> pathlib.Path:
111
+ from safetensors.torch import save_file
112
+
113
+ path = pathlib.Path(path)
114
+ path.mkdir(parents=True, exist_ok=True)
115
+ save_file({k: v.contiguous() for k, v in self.state_dict().items()}, path / "model.safetensors")
116
+ (path / "config.json").write_text(json.dumps(asdict(self.config), indent=1))
117
+ self.tokenizer.save_pretrained(path)
118
+ return path
119
+
120
+ @classmethod
121
+ def from_pretrained(cls, path: str | pathlib.Path, device: str | None = None) -> "DecisionModel":
122
+ """A local directory saved by `save`, or a model id on the Hub."""
123
+ from safetensors.torch import load_file
124
+ from transformers import AutoConfig, AutoModel, AutoTokenizer
125
+
126
+ local = pathlib.Path(path)
127
+ if not local.is_dir():
128
+ from huggingface_hub import snapshot_download
129
+
130
+ local = pathlib.Path(snapshot_download(str(path)))
131
+ config = Config(**json.loads((local / "config.json").read_text()))
132
+ tokenizer = AutoTokenizer.from_pretrained(local)
133
+ attn = "eager" if not torch.cuda.is_available() else "sdpa"
134
+ encoder = AutoModel.from_config(AutoConfig.from_pretrained(config.encoder), attn_implementation=attn)
135
+ model = cls(config=config, tokenizer=tokenizer, encoder=encoder)
136
+ model.load_state_dict(load_file(local / "model.safetensors"))
137
+ return model.to(device or ("cuda" if torch.cuda.is_available() else "cpu"))
138
+
139
+ def push_to_hub(self, repo_id: str, private: bool = False) -> str:
140
+ """Upload to the Hub. The token comes from the environment; this never asks for one."""
141
+ import tempfile
142
+
143
+ from huggingface_hub import HfApi
144
+
145
+ api = HfApi()
146
+ api.create_repo(repo_id, private=private, exist_ok=True)
147
+ with tempfile.TemporaryDirectory() as tmp:
148
+ self.save(tmp)
149
+ api.upload_folder(folder_path=tmp, repo_id=repo_id)
150
+ return f"https://huggingface.co/{repo_id}"
151
+
152
+ # ── training lives in train.py, calibration in calibrate.py ───────────────────────────────
153
+ def fit(self, examples: list[Example], **kwargs):
154
+ from .train import fit
155
+
156
+ return fit(self, examples, **kwargs)
157
+
158
+ def calibrate(self, examples: list[Example], **kwargs):
159
+ from .calibrate import calibrate
160
+
161
+ return calibrate(self, examples, **kwargs)
162
+
163
+
164
+ def encode_one(model: DecisionModel, request: Request) -> tuple[Laid, Batch]:
165
+ """Handy in tests: one request, laid out and collated."""
166
+ laid = model.layout(request.text, request.question, request.options)
167
+ return laid, collate([laid], [request], model.tokenizer.pad_token_id)