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/__init__.py +1 -0
- rulesmith/ablate.py +99 -0
- rulesmith/arena.py +172 -0
- rulesmith/bench.py +1068 -0
- rulesmith/calibrate.py +457 -0
- rulesmith/chat_judge.py +205 -0
- rulesmith/chess.py +526 -0
- rulesmith/clef.py +66 -0
- rulesmith/cli.py +1234 -0
- rulesmith/diagram.py +226 -0
- rulesmith/doom.py +550 -0
- rulesmith/extract.py +77 -0
- rulesmith/grade.py +85 -0
- rulesmith/graph.py +975 -0
- rulesmith/label.py +67 -0
- rulesmith/level.py +389 -0
- rulesmith/maps.py +96 -0
- rulesmith/mine.py +313 -0
- rulesmith/optimize.py +931 -0
- rulesmith/rules.py +1017 -0
- rulesmith/runtime.py +711 -0
- rulesmith/serve.py +68 -0
- rulesmith/tuning.py +134 -0
- rulesmith-0.1.0.dist-info/METADATA +131 -0
- rulesmith-0.1.0.dist-info/RECORD +28 -0
- rulesmith-0.1.0.dist-info/WHEEL +4 -0
- rulesmith-0.1.0.dist-info/entry_points.txt +2 -0
- rulesmith-0.1.0.dist-info/licenses/LICENSE +21 -0
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
|
+
}
|
rulesmith/chat_judge.py
ADDED
|
@@ -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()
|