rulesmith 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.
rulesmith/calibrate.py ADDED
@@ -0,0 +1,457 @@
1
+ """Fitting how sure a judge should read, and measuring how often each rule is right. A model
2
+ that says 0.97 and is right three times in five is overconfident, and a rule testing
3
+ `confidence >= 0.9` on it means nothing. One temperature per question, fitted on labeled examples,
4
+ makes its probabilities say what they claim. A rule claims nothing at all until it is measured:
5
+ counted on labeled examples, its decisions are as sure as it was right, and inputs that reach it
6
+ far more or less often than those did are no longer the inputs it was measured on. Where a rule
7
+ draws its line on a judgment, `scam >= 0.5`, is a number the proposer guessed; fitted here, it is
8
+ the line between the judge's answers that gets the most labeled examples right and holds up on
9
+ examples held out of the fit."""
10
+
11
+ import math
12
+ from collections import Counter
13
+ from dataclasses import dataclass
14
+
15
+ from rulesmith.graph import (
16
+ All,
17
+ Any,
18
+ Call,
19
+ Choice,
20
+ Coalesce,
21
+ Compare,
22
+ Condition,
23
+ Confidence,
24
+ Measured,
25
+ NodeInput,
26
+ Not,
27
+ Noul,
28
+ Plan,
29
+ Score,
30
+ Task,
31
+ Threshold,
32
+ Value,
33
+ )
34
+ from rulesmith.runtime import (
35
+ DecisionProgram,
36
+ Models,
37
+ Prices,
38
+ calibrated,
39
+ sureness,
40
+ )
41
+
42
+ # Temperatures tried, evenly spaced in log from a fifth to five: calibration corrects a judge's
43
+ # confidence, it does not invert it.
44
+ TEMPERATURES = [math.exp(math.log(0.2) + index * (math.log(25) / 80)) for index in range(81)]
45
+
46
+
47
+ def fit_temperature(pairs: list[tuple[dict[str, float], str]]) -> float:
48
+ """The temperature under which these answers give the right labels the most probability,
49
+ judged by log loss. `pairs` holds each answer's probabilities and the label that was right."""
50
+
51
+ def loss(temperature: float) -> float:
52
+ total = 0.0
53
+ for probabilities, right in pairs:
54
+ tempered = calibrated(
55
+ {"type": "choice", "value": right, "probabilities": probabilities}, temperature
56
+ )
57
+ total -= math.log(max(tempered["probabilities"][right], 1e-12))
58
+ return total
59
+
60
+ return round(min(TEMPERATURES, key=loss), 3)
61
+
62
+
63
+ def calibrate(
64
+ plan: Plan, task: Task, client, model: str, models: Models | None = None
65
+ ) -> tuple[Plan, dict[str, float]]:
66
+ """The plan with a temperature fitted for every choice question over the task's labels, from
67
+ how its answers on the examples it is fitted on compare with their labels. Other questions
68
+ have no labels to be right about and keep theirs."""
69
+ fitted_nodes = {
70
+ name: node
71
+ for name, node in plan.nodes.items()
72
+ if isinstance(node, Call)
73
+ and isinstance(node.question, Choice)
74
+ and set(node.question.criteria) <= set(task.labels)
75
+ }
76
+ raw = plan.model_copy(
77
+ update={
78
+ "nodes": plan.nodes
79
+ | {
80
+ name: node.model_copy(update={"temperature": 1})
81
+ for name, node in fitted_nodes.items()
82
+ }
83
+ }
84
+ )
85
+ program = DecisionProgram(raw, client, model, models=models)
86
+ pairs: dict[str, list] = {name: [] for name in fitted_nodes}
87
+ for example in task.fitting():
88
+ answers = program(state=example.state).answers
89
+ for name, node in fitted_nodes.items():
90
+ if name in answers and example.label in node.question.criteria:
91
+ pairs[name].append((answers[name]["probabilities"], example.label))
92
+ temperatures = {name: fit_temperature(found) for name, found in pairs.items() if found}
93
+ nodes = plan.nodes | {
94
+ name: plan.nodes[name].model_copy(update={"temperature": temperature})
95
+ for name, temperature in temperatures.items()
96
+ }
97
+ return Plan.model_validate(plan.model_copy(update={"nodes": nodes}).model_dump()), temperatures
98
+
99
+
100
+ @dataclass(frozen=True)
101
+ class Cut:
102
+ """One numeric line a rule draws on a judgment: a `>=` or `<` in a rule's condition, reached
103
+ by `path` through its and/or/not groups, or a threshold node's cutoff (`path` None). `tests`
104
+ is the judgment or confidence the line is drawn on."""
105
+
106
+ node: str
107
+ path: tuple[int, ...] | None
108
+ tests: str
109
+
110
+ def value(self, plan: Plan) -> float:
111
+ if self.path is None:
112
+ return plan.nodes[self.node].cutoff
113
+ return descend(plan.nodes[self.node].when, self.path).value
114
+
115
+ def moved(self, plan: Plan, value: float) -> Plan:
116
+ """The plan with this line drawn at `value` and nothing else changed."""
117
+ data = plan.model_dump()
118
+ node = data["nodes"][self.node]
119
+ if self.path is None:
120
+ node["cutoff"] = value
121
+ else:
122
+ condition = node["when"]
123
+ for index in self.path:
124
+ condition = (
125
+ condition["conditions"][index]
126
+ if "conditions" in condition
127
+ else condition["condition"]
128
+ )
129
+ condition["value"] = value
130
+ return Plan.model_validate(data)
131
+
132
+
133
+ def descend(condition: Condition, path: tuple[int, ...]) -> Compare:
134
+ for index in path:
135
+ condition = (
136
+ condition.conditions[index]
137
+ if isinstance(condition, (All, Any))
138
+ else condition.condition
139
+ )
140
+ return condition
141
+
142
+
143
+ def cuts(plan: Plan) -> list[Cut]:
144
+ """Every line a rule draws on a number the judge produced: a probability, a score's expected
145
+ level, or how sure a judgment was. Lines drawn on readings are what the rule says about the
146
+ input, not about the judge, and stay where they were written."""
147
+ judged = {
148
+ name
149
+ for name, node in plan.nodes.items()
150
+ if (isinstance(node, Call) and isinstance(node.question, (Noul, Score)))
151
+ or isinstance(node, Confidence)
152
+ }
153
+
154
+ def compares(condition: Condition, path: tuple[int, ...]):
155
+ if isinstance(condition, Compare):
156
+ if condition.op != "eq" and condition.node in judged:
157
+ yield path, condition.node
158
+ elif isinstance(condition, Not):
159
+ yield from compares(condition.condition, (*path, 0))
160
+ elif isinstance(condition, (All, Any)):
161
+ for index, child in enumerate(condition.conditions):
162
+ yield from compares(child, (*path, index))
163
+
164
+ found = []
165
+ for name, node in plan.nodes.items():
166
+ if isinstance(node, Threshold) and node.source in judged:
167
+ found.append(Cut(name, None, node.source))
168
+ if getattr(node, "when", None) is not None:
169
+ found += [Cut(name, path, tests) for path, tests in compares(node.when, ())]
170
+ return found
171
+
172
+
173
+ FOLDS = 5
174
+
175
+
176
+ def fit_cutoffs(
177
+ plan: Plan,
178
+ task: Task,
179
+ client,
180
+ model: str,
181
+ models: Models | None = None,
182
+ prices: Prices | None = None,
183
+ ) -> tuple[Plan, list[dict]]:
184
+ """The plan with every cut moved to where it gets the most fitting examples right, one cut
185
+ at a time, holding the questions and the rules' words still. A move is kept only when it
186
+ also wins held out: the fitting examples are dealt into folds, the best line is chosen
187
+ without each fold in turn and scored on it, and those scores together must beat the line as
188
+ it stands. One example on the far side of a line is not evidence to move it; a few are.
189
+ The report has a row per line: where it was and where it ended, how many places were tried,
190
+ the training accuracy just before and after its move, and whether it moved. A task with a
191
+ prevalence counts each example as the real traffic it stands for, so the line is placed for
192
+ the rate the graph will meet, not the rate the examples were sampled at. With `prices`, each
193
+ example also pays for the calls it made, so a line on how sure a small judge was is placed
194
+ where escalating to a dearer one pays, not wherever escalating is never worse."""
195
+ examples = task.fitting()
196
+ lines = cuts(plan)
197
+ if not examples or not lines:
198
+ return plan, []
199
+ weights = task.weights()
200
+ weight = [weights[e.label] for e in examples]
201
+ total = sum(weight)
202
+
203
+ prices = prices or Prices()
204
+
205
+ def right(candidate: Plan) -> list[float]:
206
+ decided = outcomes(candidate, examples, client, model, models)
207
+ return [
208
+ w * (record["correct"] - prices.charge(record["calls"]))
209
+ for w, record in zip(weight, decided, strict=True)
210
+ ]
211
+
212
+ seen = observed(plan, examples, client, model, models)
213
+ started = {cut: cut.value(plan) for cut in lines}
214
+ tried = dict.fromkeys(lines, 0)
215
+ after = right(plan)
216
+ change: dict[Cut, tuple[list[float], list[float]]] = {}
217
+ for _ in range(3):
218
+ moved = False
219
+ for cut in lines:
220
+ current = cut.value(plan)
221
+ candidates = midpoints(seen.get(cut.tests, []))
222
+ if not candidates:
223
+ continue
224
+ tried[cut] += len(candidates)
225
+ scores = {value: right(cut.moved(plan, value)) for value in candidates}
226
+ scores[current] = after
227
+ best = choose(scores, current)
228
+ if best != current and held_out_wins(scores, current, len(examples)):
229
+ change[cut] = (after, scores[best])
230
+ plan, after, moved = cut.moved(plan, best), scores[best], True
231
+ if not moved:
232
+ break
233
+ report = [
234
+ {
235
+ "rule": cut.node,
236
+ "tests": cut.tests,
237
+ "from": started[cut],
238
+ "to": cut.value(plan),
239
+ "tried": tried[cut],
240
+ "train_before": sum(change.get(cut, (after, after))[0]) / total,
241
+ "train_after": sum(change.get(cut, (after, after))[1]) / total,
242
+ "kept": cut.value(plan) != started[cut],
243
+ }
244
+ for cut in lines
245
+ if tried[cut]
246
+ ]
247
+ return plan, report
248
+
249
+
250
+ def observed(
251
+ plan: Plan, examples: list, client, model: str, models: Models | None = None
252
+ ) -> dict[str, list[float]]:
253
+ """The numbers each judgment and confidence took on the examples, the values a line can be
254
+ drawn between."""
255
+ program = DecisionProgram(plan, client, model, models=models)
256
+ seen: dict[str, list[float]] = {}
257
+ for example in examples:
258
+ answers = program(state=example.state).answers
259
+ for name, node in plan.nodes.items():
260
+ if (
261
+ isinstance(node, Call)
262
+ and name in answers
263
+ and isinstance(node.question, (Noul, Score))
264
+ ):
265
+ seen.setdefault(name, []).append(answers[name]["value"])
266
+ if isinstance(node, Confidence) and node.source in answers:
267
+ seen.setdefault(name, []).append(sureness(answers[node.source]))
268
+ return seen
269
+
270
+
271
+ def midpoints(values: list[float]) -> list[float]:
272
+ """A line between each pair of neighboring values, rounded where rounding still separates
273
+ them: any line between the same two answers decides the same examples, and a short number
274
+ reads as a choice rather than an accident."""
275
+ distinct = sorted(set(values))
276
+ found = []
277
+ for low, high in zip(distinct, distinct[1:], strict=False):
278
+ exact = (low + high) / 2
279
+ short = round(exact, 4)
280
+ found.append(short if low < short < high else exact)
281
+ return found
282
+
283
+
284
+ def correct(
285
+ plan: Plan, examples: list, client, model: str, models: Models | None = None
286
+ ) -> list[bool]:
287
+ return [record["correct"] for record in outcomes(plan, examples, client, model, models)]
288
+
289
+
290
+ def choose(scores: dict[float, list[float]], current: float, fold: int | None = None) -> float:
291
+ """The line with the most right answers, outside `fold` if one is named; among equals, the
292
+ one nearest to where the line stands, so a tie never moves it."""
293
+
294
+ def total(value: float) -> int:
295
+ return sum(right for index, right in enumerate(scores[value]) if index % FOLDS != fold)
296
+
297
+ return max(scores, key=lambda value: (total(value), -abs(value - current)))
298
+
299
+
300
+ def held_out_wins(scores: dict[float, list[float]], current: float, count: int) -> bool:
301
+ """Whether choosing the line without each fold and scoring it on that fold beats, over all the
302
+ folds, the line as it stands."""
303
+ folds = min(FOLDS, count)
304
+ held_out = 0
305
+ for fold in range(folds):
306
+ chosen = choose(scores, current, fold)
307
+ held_out += sum(
308
+ right for index, right in enumerate(scores[chosen]) if index % FOLDS == fold
309
+ )
310
+ return held_out > sum(scores[current])
311
+
312
+
313
+ def outcomes(
314
+ plan: Plan, examples: list, client, model: str, models: Models | None = None
315
+ ) -> list[dict]:
316
+ """Each labeled example as the plan decided it: the node that decided, the judgments asked,
317
+ and whether the answer was one the example accepts."""
318
+ program = DecisionProgram(plan, client, model, models=models)
319
+ found = []
320
+ for example in examples:
321
+ prediction = dict(program(state=example.state))
322
+ found.append(
323
+ {
324
+ "label": prediction.get("label"),
325
+ "decided_by": prediction.get("decided_by"),
326
+ "calls": prediction["calls"],
327
+ "correct": prediction.get("label") in (example.label, *example.accepted),
328
+ }
329
+ )
330
+ return found
331
+
332
+
333
+ def tally(records: list[dict]) -> dict[str | None, tuple[int, int]]:
334
+ """For each node that decided an answer, how many of its answers were right and how many it
335
+ gave. A constant graph decides with no node, under None."""
336
+ counts: dict[str | None, tuple[int, int]] = {}
337
+ for record in records:
338
+ if record.get("label") is None:
339
+ continue
340
+ right, decided = counts.get(record["decided_by"], (0, 0))
341
+ counts[record["decided_by"]] = (right + bool(record["correct"]), decided + 1)
342
+ return counts
343
+
344
+
345
+ def rules(plan: Plan) -> list[str]:
346
+ """The rules the label is decided by: every value or threshold in its list of fallbacks."""
347
+ source = plan.outputs["label"]
348
+ pending, found = [source.node] if isinstance(source, NodeInput) else [], []
349
+ while pending:
350
+ name = pending.pop(0)
351
+ node = plan.nodes[name]
352
+ if isinstance(node, Coalesce):
353
+ pending = list(node.sources) + pending
354
+ elif isinstance(node, (Value, Threshold)):
355
+ found.append(name)
356
+ return found
357
+
358
+
359
+ def measure(plan: Plan, task: Task, client, model: str, models: Models | None = None) -> Plan:
360
+ """The plan with every rule that decides its label counted on the task's calibration
361
+ examples, or without those on its labeled training and validation examples. A rule decides
362
+ only some of them, so it is counted on all it can be; the training and validation examples
363
+ are ones search has seen, so those counts flatter a searched graph."""
364
+ examples = task.fitting()
365
+ if not task.calibration:
366
+ examples += [e for e in task.validation if e.label is not None]
367
+ if not examples or not rules(plan):
368
+ return plan
369
+ counts = tally(outcomes(plan, examples, client, model, models))
370
+ nodes = plan.nodes | {
371
+ name: plan.nodes[name].model_copy(
372
+ update={
373
+ "measured": Measured(
374
+ right=counts.get(name, (0, 0))[0],
375
+ decided=counts.get(name, (0, 0))[1],
376
+ examples=len(examples),
377
+ )
378
+ }
379
+ )
380
+ for name in rules(plan)
381
+ }
382
+ return Plan.model_validate(plan.model_copy(update={"nodes": nodes}).model_dump())
383
+
384
+
385
+ def unmeasured(plan: Plan) -> Plan:
386
+ """The plan without its rules' counts, which describe this graph and no rewrite of it."""
387
+ nodes = {
388
+ name: node.model_copy(update={"measured": None})
389
+ if isinstance(node, (Value, Threshold))
390
+ else node
391
+ for name, node in plan.nodes.items()
392
+ }
393
+ return plan.model_copy(update={"nodes": nodes})
394
+
395
+
396
+ def binomial_p(count: int, trials: int, share: float) -> float:
397
+ """Exact two-sided binomial test: the chance, over `trials` examples each reaching a rule
398
+ with probability `share`, of a count no likelier than `count`."""
399
+ if share in (0, 1):
400
+ return float(count == trials * share)
401
+
402
+ def log_chance(k: int) -> float:
403
+ ways = math.lgamma(trials + 1) - math.lgamma(k + 1) - math.lgamma(trials - k + 1)
404
+ return ways + k * math.log(share) + (trials - k) * math.log1p(-share)
405
+
406
+ # A relative tolerance, as exact tests use, so a tie with the observed count is not lost
407
+ # to rounding.
408
+ observed = log_chance(count) + 1e-7
409
+ return min(1.0, sum(math.exp(c) for k in range(trials + 1) if (c := log_chance(k)) <= observed))
410
+
411
+
412
+ def drift(plan: Plan, records: list[dict], alpha: float) -> list[dict]:
413
+ """Every measured rule these examples reach at a rate its measurement makes unlikely, below
414
+ `alpha`: the inputs have moved from the ones it was counted on, and so may how often it is
415
+ right."""
416
+ reached = Counter(record.get("decided_by") for record in records)
417
+ flagged = []
418
+ for name, node in plan.nodes.items():
419
+ measured = getattr(node, "measured", None)
420
+ if measured is None:
421
+ continue
422
+ p = binomial_p(reached[name], len(records), measured.share)
423
+ if p < alpha:
424
+ flagged.append(
425
+ {
426
+ "rule": name,
427
+ "measured_share": measured.share,
428
+ "share": reached[name] / len(records),
429
+ "p": p,
430
+ }
431
+ )
432
+ return flagged
433
+
434
+
435
+ def at_prevalence(task: Task, records: list[dict]) -> dict | None:
436
+ """Accuracy, and each label's precision and recall, as they would be on real traffic: each
437
+ scored example counted as the share of it that its label stands for. None for a task with no
438
+ prevalence, whose examples are taken to be sampled as they come."""
439
+ if task.prevalence is None:
440
+ return None
441
+ weights = task.weights()
442
+ scored = [(r["expected"], r.get("label"), weights[r["expected"]]) for r in records]
443
+ total = sum(w for _, _, w in scored)
444
+ labels = {}
445
+ for label in task.labels:
446
+ hit = sum(w for e, a, w in scored if e == label and a == label)
447
+ answered = sum(w for _, a, w in scored if a == label)
448
+ present = sum(w for e, _, w in scored if e == label)
449
+ labels[label] = {
450
+ "precision": hit / answered if answered else None,
451
+ "recall": hit / present if present else None,
452
+ }
453
+ return {
454
+ "prevalence": task.prevalence,
455
+ "accuracy": sum(w for e, a, w in scored if e == a) / total,
456
+ "labels": labels,
457
+ }
@@ -0,0 +1,205 @@
1
+ """A judge reached through any OpenAI-compatible chat API, read from its log-probabilities.
2
+
3
+ Each question is asked on its own, with its options lettered A, B, C and so on, for a single
4
+ token. The probability of each option is the softmax of its letter's log-probability among the
5
+ first token's most likely candidates, so an answer carries the same distribution a TypeSafe
6
+ judge returns, and the runtime reads it the same way.
7
+ """
8
+
9
+ import json
10
+ import math
11
+ import string
12
+ import time
13
+ from types import TracebackType
14
+ from typing import Self
15
+
16
+ import httpx
17
+ from typesafe_sdk import (
18
+ ChoiceAnswer,
19
+ NoulAnswer,
20
+ RetryPolicy,
21
+ ScoreAnswer,
22
+ SystemOneResponse,
23
+ Usage,
24
+ )
25
+
26
+ LETTERS = string.ascii_uppercase
27
+ # How many first-token candidates to ask for. OpenAI and vLLM allow 20, but mlx_lm.server
28
+ # refuses more than 11 by dropping the connection; ten covers every score rubric and most
29
+ # choices, and a letter beyond them is bounded by ABSENT_MARGIN.
30
+ TOP_LOGPROBS = 10
31
+ # A letter missing from the top candidates is less likely than every listed one by an unknown
32
+ # amount. It is taken to sit this many nats below the least likely listed token: low enough to
33
+ # rank under all of them, while still keeping a little probability, as the model would.
34
+ ABSENT_MARGIN = 1.0
35
+ SYSTEM = "You are a careful judge. Answer with exactly one letter, and nothing else."
36
+
37
+
38
+ class ChatJudgeError(Exception):
39
+ """The chat API failed, or answered in a form a judgment cannot be read from."""
40
+
41
+
42
+ def letter_probabilities(candidates: list[dict], count: int) -> list[float]:
43
+ """The probabilities of the first `count` letters, from the first token's top candidates."""
44
+ letters = LETTERS[:count]
45
+ found: dict[str, list[float]] = {}
46
+ for candidate in candidates:
47
+ # Some servers return raw tokenizer pieces, with the leading space spelled as byte-level
48
+ # BPE's "Ġ" or SentencePiece's "▁".
49
+ letter = candidate["token"].strip().lstrip("Ġ▁")
50
+ if letter in letters:
51
+ found.setdefault(letter, []).append(candidate["logprob"])
52
+ if not found:
53
+ likeliest = max(candidates, key=lambda candidate: candidate["logprob"])["token"]
54
+ raise ChatJudgeError(
55
+ "none of the option letters is among the model's most likely first tokens; it "
56
+ f"began with {likeliest!r}, as a model that thinks before answering does"
57
+ )
58
+ floor = min(candidate["logprob"] for candidate in candidates) - ABSENT_MARGIN
59
+ logits = [logsumexp(found[letter]) if letter in found else floor for letter in letters]
60
+ peak = max(logits)
61
+ weights = [math.exp(logit - peak) for logit in logits]
62
+ total = sum(weights)
63
+ return [weight / total for weight in weights]
64
+
65
+
66
+ def logsumexp(values: list[float]) -> float:
67
+ peak = max(values)
68
+ return peak + math.log(sum(math.exp(value - peak) for value in values))
69
+
70
+
71
+ # Confidence follows the formulas TypeSafe publishes in its MIT-licensed system-one-adapter (as
72
+ # Ruling's calibration.py reimplements them), so it reads the same as a hosted judge's.
73
+ def choice_confidence(probabilities: list[float]) -> float:
74
+ """How far the leading option stands above a uniform distribution, from 0 to 1."""
75
+ uniform = 1 / len(probabilities)
76
+ return (max(probabilities) - uniform) / (1 - uniform)
77
+
78
+
79
+ def score_confidence(probabilities: list[float]) -> float:
80
+ """How tightly a score's mass sits around its modal level, from 0 to 1."""
81
+ mode = probabilities.index(max(probabilities))
82
+ spread = sum(p * abs(level - mode) for level, p in enumerate(probabilities))
83
+ middle = (len(probabilities) - 1) / 2
84
+ uniform_spread = sum(abs(level - middle) for level in range(len(probabilities)))
85
+ return max(0.0, 1 - spread * len(probabilities) / uniform_spread)
86
+
87
+
88
+ def options(question: dict) -> list[str]:
89
+ """The options a question offers, in the order they are lettered."""
90
+ if question["type"] == "noul":
91
+ return ["True", "False"]
92
+ if question["type"] == "choice":
93
+ return [f"{label}: {meaning}" for label, meaning in question["criteria"].items()]
94
+ return list(question["criteria"])
95
+
96
+
97
+ def prompt(state, question: dict) -> str:
98
+ shown = state if isinstance(state, str) else json.dumps(state, indent=2)
99
+ lettered = "\n".join(f"{LETTERS[i]}. {text}" for i, text in enumerate(options(question)))
100
+ heading = {
101
+ "noul": "Is the statement true of the input?",
102
+ "choice": "Options:",
103
+ "score": "Levels, from lowest to highest:",
104
+ }[question["type"]]
105
+ return (
106
+ f"Input:\n{shown}\n\n{question['instructions']}\n\n{heading}\n{lettered}\n\n"
107
+ "Answer with exactly one letter."
108
+ )
109
+
110
+
111
+ def answer(question: dict, probabilities: list[float]):
112
+ if question["type"] == "noul":
113
+ return NoulAnswer(noul=probabilities[0])
114
+ if question["type"] == "choice":
115
+ labels = list(question["criteria"])
116
+ return ChoiceAnswer(
117
+ choice=labels[probabilities.index(max(probabilities))],
118
+ confidence=choice_confidence(probabilities),
119
+ probabilities=dict(zip(labels, probabilities, strict=True)),
120
+ )
121
+ return ScoreAnswer(
122
+ score=sum(level * p for level, p in enumerate(probabilities)),
123
+ confidence=score_confidence(probabilities),
124
+ legend=dict(enumerate(question["criteria"])),
125
+ probabilities=dict(enumerate(probabilities)),
126
+ )
127
+
128
+
129
+ class ChatClient:
130
+ """Answers TypeSafe's System One questions through an OpenAI-compatible chat API."""
131
+
132
+ def __init__(self, *, base_url: str, api_key: str | None, timeout: float, retry: RetryPolicy):
133
+ # A local server needs no key, and an absent one is better than a made-up one.
134
+ headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
135
+ self.http = httpx.Client(base_url=base_url.rstrip("/"), headers=headers, timeout=timeout)
136
+ self.retry = retry
137
+
138
+ def system_one(self, *, state, questions: dict[str, dict], model: str) -> SystemOneResponse:
139
+ answers, served, input_tokens, output_tokens = {}, model, 0, 0
140
+ for key, question in questions.items():
141
+ count = len(options(question))
142
+ if count > len(LETTERS):
143
+ raise ChatJudgeError(
144
+ f"question {key!r} has {count} options, and the chat judge letters at most "
145
+ f"{len(LETTERS)}"
146
+ )
147
+ body = self.complete(model, prompt(state, question))
148
+ choice = body["choices"][0]
149
+ content = (choice.get("logprobs") or {}).get("content")
150
+ if not content or not content[0].get("top_logprobs"):
151
+ raise ChatJudgeError(
152
+ f"{model} returned no log-probabilities; the openai backend needs a model "
153
+ "that returns them, which reasoning models and some providers do not"
154
+ )
155
+ answers[key] = answer(question, letter_probabilities(content[0]["top_logprobs"], count))
156
+ served = body.get("model", served)
157
+ usage = body.get("usage") or {}
158
+ input_tokens += usage.get("prompt_tokens", 0)
159
+ output_tokens += usage.get("completion_tokens", 0)
160
+ return SystemOneResponse(
161
+ model=served,
162
+ usage=Usage(input_tokens=input_tokens, output_tokens=output_tokens),
163
+ answers=answers,
164
+ )
165
+
166
+ def complete(self, model: str, text: str) -> dict:
167
+ request = {
168
+ "model": model,
169
+ "messages": [{"role": "system", "content": SYSTEM}, {"role": "user", "content": text}],
170
+ "max_tokens": 1,
171
+ "temperature": 0,
172
+ "logprobs": True,
173
+ "top_logprobs": TOP_LOGPROBS,
174
+ }
175
+ attempt = 0
176
+ while True:
177
+ last = attempt == self.retry.max_retries
178
+ try:
179
+ response = self.http.post("/chat/completions", json=request)
180
+ except httpx.TransportError as error:
181
+ if last:
182
+ raise ChatJudgeError(f"{self.http.base_url} is unreachable: {error}") from error
183
+ else:
184
+ if response.is_success:
185
+ return response.json()
186
+ if last or response.status_code not in self.retry.http_statuses:
187
+ raise ChatJudgeError(
188
+ f"{self.http.base_url} answered {response.status_code}: {response.text}"
189
+ )
190
+ time.sleep(min(self.retry.backoff_initial * 2**attempt, self.retry.backoff_max))
191
+ attempt += 1
192
+
193
+ def close(self) -> None:
194
+ self.http.close()
195
+
196
+ def __enter__(self) -> Self:
197
+ return self
198
+
199
+ def __exit__(
200
+ self,
201
+ exc_type: type[BaseException] | None,
202
+ exc_value: BaseException | None,
203
+ traceback: TracebackType | None,
204
+ ) -> None:
205
+ self.close()