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/cli.py
ADDED
|
@@ -0,0 +1,1234 @@
|
|
|
1
|
+
"""Validate, execute, evaluate, and evolve a saved decision graph."""
|
|
2
|
+
|
|
3
|
+
import argparse
|
|
4
|
+
import importlib
|
|
5
|
+
import json
|
|
6
|
+
import os
|
|
7
|
+
import signal
|
|
8
|
+
import subprocess
|
|
9
|
+
import sys
|
|
10
|
+
from collections import Counter
|
|
11
|
+
from contextlib import ExitStack, nullcontext, redirect_stdout
|
|
12
|
+
from dataclasses import dataclass
|
|
13
|
+
from pathlib import Path
|
|
14
|
+
from types import ModuleType
|
|
15
|
+
from typing import TYPE_CHECKING
|
|
16
|
+
from urllib.parse import urlsplit
|
|
17
|
+
|
|
18
|
+
import dspy
|
|
19
|
+
from typesafe_sdk import RetryPolicy, TypeSafeClient, TypeSafeError
|
|
20
|
+
|
|
21
|
+
from rulesmith import clef
|
|
22
|
+
from rulesmith.ablate import effects
|
|
23
|
+
from rulesmith.calibrate import at_prevalence, calibrate, correct, drift, fit_cutoffs, measure
|
|
24
|
+
from rulesmith.chat_judge import ChatClient, ChatJudgeError
|
|
25
|
+
from rulesmith.diagram import mermaid
|
|
26
|
+
from rulesmith.extract import Extractor
|
|
27
|
+
from rulesmith.grade import Graded, Grader, model_name
|
|
28
|
+
from rulesmith.graph import Graph, Plan, Task
|
|
29
|
+
from rulesmith.label import Labeler, Unlabeled, draft_task
|
|
30
|
+
from rulesmith.maps import MapError, load, unit_task, units
|
|
31
|
+
from rulesmith.optimize import (
|
|
32
|
+
GraphProposer,
|
|
33
|
+
Labels,
|
|
34
|
+
SearchConfig,
|
|
35
|
+
described_labels,
|
|
36
|
+
evaluate_plan,
|
|
37
|
+
optimize_plan,
|
|
38
|
+
)
|
|
39
|
+
from rulesmith.rules import NAME, RuleError, keep_comments, render_rules
|
|
40
|
+
from rulesmith.runtime import (
|
|
41
|
+
DecisionProgram,
|
|
42
|
+
ExecutionConfig,
|
|
43
|
+
Judge,
|
|
44
|
+
Models,
|
|
45
|
+
Prices,
|
|
46
|
+
Remembered,
|
|
47
|
+
digest,
|
|
48
|
+
)
|
|
49
|
+
from rulesmith.serve import make_server
|
|
50
|
+
from rulesmith.tuning import BASE, FILES, taught, tune
|
|
51
|
+
|
|
52
|
+
# The games load only for a task that plays one (see game), so a core install without their
|
|
53
|
+
# packages still validates and runs every other graph.
|
|
54
|
+
if TYPE_CHECKING:
|
|
55
|
+
from rulesmith.chess import ChessTask
|
|
56
|
+
from rulesmith.doom import DoomTask
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
@dataclass(frozen=True)
|
|
60
|
+
class Backend:
|
|
61
|
+
model: str
|
|
62
|
+
base_url_env: str
|
|
63
|
+
api_key_env: str
|
|
64
|
+
base_url: str | None = None
|
|
65
|
+
api_key: str | None = None
|
|
66
|
+
client: type[TypeSafeClient] | type[ChatClient] = TypeSafeClient
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
BACKENDS = {
|
|
70
|
+
"typesafe": Backend("jev-latest", "TYPESAFE_BASE_URL", "TYPESAFE_API_KEY"),
|
|
71
|
+
# The SDK requires a nonempty key even when the local server disables auth.
|
|
72
|
+
"ruling": Backend(
|
|
73
|
+
"default", "RULING_BASE_URL", "RULING_API_KEY", "http://127.0.0.1:8010", "local-no-auth"
|
|
74
|
+
),
|
|
75
|
+
# Cloudflare's Clef-flash, served on this machine by `rulesmith clef-serve` or by the MLX
|
|
76
|
+
# port's own clef_mlx.py, which answer the same /v1/systemone requests, one at a time.
|
|
77
|
+
"clef": Backend(
|
|
78
|
+
"clef-flash", "CLEF_BASE_URL", "CLEF_API_KEY", "http://127.0.0.1:8000", "local-no-auth"
|
|
79
|
+
),
|
|
80
|
+
# Any OpenAI-compatible chat API whose models return log-probabilities, read as a judge.
|
|
81
|
+
# OpenAI's newer GPT-5 and GPT-6 models return none, so the default is gpt-4.1.
|
|
82
|
+
"openai": Backend(
|
|
83
|
+
"gpt-4.1",
|
|
84
|
+
"OPENAI_BASE_URL",
|
|
85
|
+
"OPENAI_API_KEY",
|
|
86
|
+
"https://api.openai.com/v1",
|
|
87
|
+
client=ChatClient,
|
|
88
|
+
),
|
|
89
|
+
}
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
def resolve_backend(name: str, model: str | None = None, base_url: str | None = None):
|
|
93
|
+
backend = BACKENDS[name]
|
|
94
|
+
endpoint = (base_url or os.environ.get(backend.base_url_env, "")).strip() or backend.base_url
|
|
95
|
+
key = os.environ.get(backend.api_key_env, "").strip() or backend.api_key
|
|
96
|
+
return model or backend.model, {"base_url": endpoint, "api_key": key}
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
def connect(name: str, connection: dict, timeout: float, retry: RetryPolicy):
|
|
100
|
+
"""A client for the backend `name`, in the protocol that backend speaks."""
|
|
101
|
+
return BACKENDS[name].client(**connection, timeout=timeout, retry=retry)
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
def parser() -> argparse.ArgumentParser:
|
|
105
|
+
root = argparse.ArgumentParser(description=__doc__)
|
|
106
|
+
commands = root.add_subparsers(dest="command", required=True)
|
|
107
|
+
helps = {
|
|
108
|
+
"validate": "Check a graph against a task without calling any model",
|
|
109
|
+
"evaluate": "Score a graph on one split of a task",
|
|
110
|
+
"optimize": "Search for the graph with the best validation score",
|
|
111
|
+
"record": "Play a graph in Doom and save the video",
|
|
112
|
+
}
|
|
113
|
+
for name, summary in helps.items():
|
|
114
|
+
command = commands.add_parser(name, help=summary)
|
|
115
|
+
command.add_argument("task", type=Path)
|
|
116
|
+
if name == "optimize":
|
|
117
|
+
command.add_argument(
|
|
118
|
+
"plan", type=Path, nargs="?", help="Seed graph; defaults to the majority label"
|
|
119
|
+
)
|
|
120
|
+
else:
|
|
121
|
+
command.add_argument("plan", type=Path)
|
|
122
|
+
command.add_argument("--max-nodes", type=int, default=SearchConfig().max_nodes)
|
|
123
|
+
command.add_argument("--max-depth", type=int, default=SearchConfig().max_depth)
|
|
124
|
+
if name == "validate":
|
|
125
|
+
continue
|
|
126
|
+
provider_arguments(command)
|
|
127
|
+
command.add_argument("--output", type=Path, required=True)
|
|
128
|
+
if name in ("evaluate", "optimize"):
|
|
129
|
+
command.add_argument(
|
|
130
|
+
"--parallel-examples",
|
|
131
|
+
type=positive,
|
|
132
|
+
default=SearchConfig().parallel_examples,
|
|
133
|
+
help="Examples or episodes scored at once; deathmatch episodes can run together",
|
|
134
|
+
)
|
|
135
|
+
if name in ("evaluate", "optimize", "record"):
|
|
136
|
+
command.add_argument(
|
|
137
|
+
"--opponent",
|
|
138
|
+
type=Path,
|
|
139
|
+
help="Graph a deathmatch task's episodes are played against",
|
|
140
|
+
)
|
|
141
|
+
if name in ("evaluate", "optimize"):
|
|
142
|
+
command.add_argument(
|
|
143
|
+
"--engine",
|
|
144
|
+
help="A UCI engine, such as stockfish, to play a chess task against instead of "
|
|
145
|
+
"a graph",
|
|
146
|
+
)
|
|
147
|
+
command.add_argument(
|
|
148
|
+
"--engine-skill", type=int, help="The engine's Skill Level (default 0, weakest)"
|
|
149
|
+
)
|
|
150
|
+
command.add_argument(
|
|
151
|
+
"--engine-seconds", type=float, help="The engine's time per move (default 0.05)"
|
|
152
|
+
)
|
|
153
|
+
command.add_argument(
|
|
154
|
+
"--grader-model",
|
|
155
|
+
help="DSPy model identifier that grades answers instead of the task's labels",
|
|
156
|
+
)
|
|
157
|
+
command.add_argument("--grader-base-url")
|
|
158
|
+
if name == "record":
|
|
159
|
+
command.add_argument(
|
|
160
|
+
"--seed", type=int, nargs="+", required=True, help="Episode seeds, played in order"
|
|
161
|
+
)
|
|
162
|
+
elif name == "evaluate":
|
|
163
|
+
command.add_argument(
|
|
164
|
+
"--split", choices=("train", "validation", "calibration", "test"), default="test"
|
|
165
|
+
)
|
|
166
|
+
command.add_argument(
|
|
167
|
+
"--drift",
|
|
168
|
+
type=float,
|
|
169
|
+
default=0.01,
|
|
170
|
+
help="Flag a calibrated rule that decides at a rate this unlikely for it",
|
|
171
|
+
)
|
|
172
|
+
command.add_argument(
|
|
173
|
+
"--visible", action="store_true", help="Show the game window for a Doom task"
|
|
174
|
+
)
|
|
175
|
+
else:
|
|
176
|
+
command.add_argument(
|
|
177
|
+
"--reflection-model",
|
|
178
|
+
required=True,
|
|
179
|
+
help="DSPy model identifier for the generative proposer",
|
|
180
|
+
)
|
|
181
|
+
command.add_argument("--reflection-base-url")
|
|
182
|
+
command.add_argument("--reflection-max-tokens", type=int, default=4096)
|
|
183
|
+
command.add_argument(
|
|
184
|
+
"--reflection-options",
|
|
185
|
+
type=json.loads,
|
|
186
|
+
help="Extra request JSON for the proposer, such as a cap on a thinking model's "
|
|
187
|
+
'reasoning: \'{"reasoning": {"max_tokens": 2048}}\'',
|
|
188
|
+
)
|
|
189
|
+
command.add_argument(
|
|
190
|
+
"--reflection-timeout", type=float, help="Seconds before a proposal request fails"
|
|
191
|
+
)
|
|
192
|
+
command.add_argument(
|
|
193
|
+
"--max-metric-calls", type=int, default=SearchConfig().max_metric_calls
|
|
194
|
+
)
|
|
195
|
+
command.add_argument(
|
|
196
|
+
"--reflection-batch", type=int, default=SearchConfig().reflection_batch
|
|
197
|
+
)
|
|
198
|
+
command.add_argument(
|
|
199
|
+
"--repair-attempts", type=int, default=SearchConfig().repair_attempts
|
|
200
|
+
)
|
|
201
|
+
command.add_argument("--merges", type=int, default=SearchConfig().merges)
|
|
202
|
+
command.add_argument(
|
|
203
|
+
"--mine",
|
|
204
|
+
action="store_true",
|
|
205
|
+
help="Mine rules a field alone settles, and start there if they score better",
|
|
206
|
+
)
|
|
207
|
+
command.add_argument(
|
|
208
|
+
"--evidence",
|
|
209
|
+
action="store_true",
|
|
210
|
+
help="Show the proposer the single-field rules the training examples support "
|
|
211
|
+
"best; counted on a sampled training set, they can over-fit its rate",
|
|
212
|
+
)
|
|
213
|
+
command.add_argument(
|
|
214
|
+
"--significance",
|
|
215
|
+
type=float,
|
|
216
|
+
help="Keep the seed unless the winner's gain passes a sign test at this level",
|
|
217
|
+
)
|
|
218
|
+
command.add_argument(
|
|
219
|
+
"--resume",
|
|
220
|
+
action="store_true",
|
|
221
|
+
help="Continue the search checkpointed under --output/search, to "
|
|
222
|
+
"--max-metric-calls; the same task and starting graph are required",
|
|
223
|
+
)
|
|
224
|
+
command.add_argument(
|
|
225
|
+
"--unit",
|
|
226
|
+
help="Search this unit of the map given as the plan, against the task given, "
|
|
227
|
+
"with the rest of the map frozen; the winner must still fit the map",
|
|
228
|
+
)
|
|
229
|
+
command.add_argument(
|
|
230
|
+
"--map-task",
|
|
231
|
+
type=Path,
|
|
232
|
+
help="With --unit: a task for the whole map, scored on its validation examples "
|
|
233
|
+
"before and after the unit changes; a winner that makes the map worse is refused",
|
|
234
|
+
)
|
|
235
|
+
cost_arguments(command)
|
|
236
|
+
command.add_argument(
|
|
237
|
+
"--keep-judge",
|
|
238
|
+
type=float,
|
|
239
|
+
help="Let rules answer without the judge only where they are right at least this "
|
|
240
|
+
"often on the labeled examples, such as 0.99",
|
|
241
|
+
)
|
|
242
|
+
command.add_argument("--seed", type=int, default=SearchConfig().seed)
|
|
243
|
+
draft = commands.add_parser("draft-labels", help="Label raw inputs for review as a task")
|
|
244
|
+
draft.add_argument("inputs", type=Path, help="JSON with description, labels, and inputs")
|
|
245
|
+
draft.add_argument("--output", type=Path, required=True)
|
|
246
|
+
draft.add_argument(
|
|
247
|
+
"--labeler-model", required=True, help="DSPy model identifier for the labeler"
|
|
248
|
+
)
|
|
249
|
+
draft.add_argument("--labeler-base-url")
|
|
250
|
+
draft.add_argument("--validation", type=int, required=True, help="Validation example count")
|
|
251
|
+
draft.add_argument("--test", type=int, required=True, help="Test example count")
|
|
252
|
+
draft.add_argument("--seed", type=int, default=0, help="Split shuffle seed")
|
|
253
|
+
arena = commands.add_parser("arena", help="Rank Doom graphs by Elo from deathmatch duels")
|
|
254
|
+
arena.add_argument("task", type=Path, help="Doom task whose episodes are the matches")
|
|
255
|
+
arena.add_argument("plans", type=Path, nargs="+", help="Entrants, named by file stem")
|
|
256
|
+
arena.add_argument("--split", choices=("train", "validation", "test"), default="test")
|
|
257
|
+
arena.add_argument("--max-nodes", type=int, default=SearchConfig().max_nodes)
|
|
258
|
+
arena.add_argument("--max-depth", type=int, default=SearchConfig().max_depth)
|
|
259
|
+
arena.add_argument("--output", type=Path, required=True)
|
|
260
|
+
provider_arguments(arena)
|
|
261
|
+
fit = commands.add_parser(
|
|
262
|
+
"calibrate",
|
|
263
|
+
help="Fit how sure each judgment reads and where rules draw lines on it, from the labels",
|
|
264
|
+
)
|
|
265
|
+
fit.add_argument("task", type=Path)
|
|
266
|
+
fit.add_argument("plan", type=Path)
|
|
267
|
+
fit.add_argument("--output", type=Path, required=True, help="A .rules or .json plan")
|
|
268
|
+
provider_arguments(fit)
|
|
269
|
+
cost_arguments(fit)
|
|
270
|
+
unit = commands.add_parser(
|
|
271
|
+
"unit-task",
|
|
272
|
+
help="Write the task of one unit of a map: the cases the map hands it, as handed",
|
|
273
|
+
)
|
|
274
|
+
unit.add_argument("task", type=Path, help="A labeled task for the whole map")
|
|
275
|
+
unit.add_argument("plan", type=Path, help="The map")
|
|
276
|
+
unit.add_argument("--unit", required=True, help="The unit's node in the map")
|
|
277
|
+
unit.add_argument("--output", type=Path, required=True)
|
|
278
|
+
provider_arguments(unit)
|
|
279
|
+
draw = commands.add_parser("draw", help="Write a graph as a Mermaid flowchart")
|
|
280
|
+
draw.add_argument("plan", type=Path)
|
|
281
|
+
draw.add_argument(
|
|
282
|
+
"--closed",
|
|
283
|
+
action="store_true",
|
|
284
|
+
help="Draw each unit as one box rather than opening it, so a map reads as a map",
|
|
285
|
+
)
|
|
286
|
+
for name, summary in (
|
|
287
|
+
("tune", "Fine-tune the judge on a task's examples with Ruling, test split held out"),
|
|
288
|
+
(
|
|
289
|
+
"distill",
|
|
290
|
+
"Fine-tune the judge on a larger model's answers to a task's inputs, labeled or not",
|
|
291
|
+
),
|
|
292
|
+
):
|
|
293
|
+
tuning = commands.add_parser(name, help=summary)
|
|
294
|
+
tuning.add_argument("task", type=Path)
|
|
295
|
+
tuning.add_argument(
|
|
296
|
+
"--output", type=Path, required=True, help="Directory for the records and the adapter"
|
|
297
|
+
)
|
|
298
|
+
tuning.add_argument(
|
|
299
|
+
"--ruling-dir",
|
|
300
|
+
type=Path,
|
|
301
|
+
default=Path(os.environ.get("RULING_DIR", "../ruling")),
|
|
302
|
+
help="A Ruling checkout (default: $RULING_DIR, else ../ruling)",
|
|
303
|
+
)
|
|
304
|
+
tuning.add_argument("--base", default=BASE, help="The model the adapter is trained on")
|
|
305
|
+
if name == "distill":
|
|
306
|
+
tuning.add_argument(
|
|
307
|
+
"--teacher",
|
|
308
|
+
type=judge_spec,
|
|
309
|
+
required=True,
|
|
310
|
+
metavar="BACKEND[:MODEL][@URL]",
|
|
311
|
+
help="The judge whose answers the student learns, such as "
|
|
312
|
+
"clef or ruling:qwen-35b@http://127.0.0.1:8012",
|
|
313
|
+
)
|
|
314
|
+
tuning.add_argument("--timeout", type=float, default=60)
|
|
315
|
+
tuning.epilog = "Anything after -- goes to `ruling train` and overrides the recipe."
|
|
316
|
+
ablation = commands.add_parser(
|
|
317
|
+
"ablate", help="Measure what each reading, judgment and name in a graph is worth"
|
|
318
|
+
)
|
|
319
|
+
ablation.add_argument("task", type=Path)
|
|
320
|
+
ablation.add_argument("plan", type=Path)
|
|
321
|
+
ablation.add_argument("--split", choices=("train", "validation", "test"), default="validation")
|
|
322
|
+
ablation.add_argument("--output", type=Path)
|
|
323
|
+
provider_arguments(ablation)
|
|
324
|
+
watching = commands.add_parser(
|
|
325
|
+
"drift",
|
|
326
|
+
help="Check a decision's traces against what its rules were calibrated on; exits 1 when a "
|
|
327
|
+
"rule has drifted",
|
|
328
|
+
)
|
|
329
|
+
watching.add_argument("decision", type=Path, help="A decision folder or graph file")
|
|
330
|
+
watching.add_argument(
|
|
331
|
+
"traces", type=Path, nargs="+", help="Trace files, or directories of them"
|
|
332
|
+
)
|
|
333
|
+
watching.add_argument(
|
|
334
|
+
"--significance",
|
|
335
|
+
type=float,
|
|
336
|
+
default=0.01,
|
|
337
|
+
help="Flag a rule whose share of decisions is this unlikely under its calibrated share",
|
|
338
|
+
)
|
|
339
|
+
serving = commands.add_parser(
|
|
340
|
+
"serve", help="Answer decisions over HTTP at POST /decisions/<name>, as run prints them"
|
|
341
|
+
)
|
|
342
|
+
serving.add_argument(
|
|
343
|
+
"decisions",
|
|
344
|
+
type=Path,
|
|
345
|
+
nargs="+",
|
|
346
|
+
help="Decision folders, each with a graph.rules or map.json, or graph files",
|
|
347
|
+
)
|
|
348
|
+
serving.add_argument("--host", default="127.0.0.1", help="The address to listen on")
|
|
349
|
+
serving.add_argument("--port", type=int, default=8020, help="0 picks a free port")
|
|
350
|
+
serving.add_argument("--trace", type=Path, help="Append every decision to this JSON-lines file")
|
|
351
|
+
provider_arguments(serving)
|
|
352
|
+
clef_serving = commands.add_parser(
|
|
353
|
+
"clef-serve",
|
|
354
|
+
help="Serve Cloudflare's Clef at POST /v1/systemone for --backend clef, with PyTorch on "
|
|
355
|
+
"an NVIDIA GPU, Apple's, or a CPU (needs the clef extra)",
|
|
356
|
+
)
|
|
357
|
+
clef_serving.add_argument(
|
|
358
|
+
"--weights",
|
|
359
|
+
default="Cloudflare/clef-flash",
|
|
360
|
+
help="A Clef release on Hugging Face, such as Cloudflare/clef for the 27B, or a "
|
|
361
|
+
"downloaded copy of one (default: Cloudflare/clef-flash)",
|
|
362
|
+
)
|
|
363
|
+
clef_serving.add_argument("--device", help="cuda, mps or cpu (default: the best there is)")
|
|
364
|
+
clef_serving.add_argument("--host", default="127.0.0.1", help="The address to listen on")
|
|
365
|
+
clef_serving.add_argument(
|
|
366
|
+
"--port",
|
|
367
|
+
type=int,
|
|
368
|
+
default=urlsplit(BACKENDS["clef"].base_url).port,
|
|
369
|
+
help="0 picks a free port (default: the one --backend clef asks)",
|
|
370
|
+
)
|
|
371
|
+
run = commands.add_parser("run", help="Run a graph on one input and print what it decided")
|
|
372
|
+
run.add_argument("plan", type=Path)
|
|
373
|
+
run.add_argument("state", type=Path, help="JSON file containing the input state")
|
|
374
|
+
run.add_argument("--trace", type=Path, help="Append the decision to this JSON-lines file")
|
|
375
|
+
provider_arguments(run)
|
|
376
|
+
return root
|
|
377
|
+
|
|
378
|
+
|
|
379
|
+
def provider_arguments(command: argparse.ArgumentParser) -> None:
|
|
380
|
+
command.add_argument("--backend", choices=tuple(BACKENDS), default="typesafe")
|
|
381
|
+
command.add_argument("--model", help="Model ID; defaults to the selected backend alias")
|
|
382
|
+
command.add_argument("--base-url", help="Override the selected backend endpoint")
|
|
383
|
+
command.add_argument("--timeout", type=float, default=60)
|
|
384
|
+
command.add_argument(
|
|
385
|
+
"--retries", type=int, default=0, help="Retries for failed or timed-out judge requests"
|
|
386
|
+
)
|
|
387
|
+
command.add_argument("--max-workers", type=positive, default=ExecutionConfig().max_workers)
|
|
388
|
+
command.add_argument(
|
|
389
|
+
"--judge",
|
|
390
|
+
type=named_judge,
|
|
391
|
+
action="append",
|
|
392
|
+
default=[],
|
|
393
|
+
metavar="NAME=BACKEND[:MODEL][@URL]",
|
|
394
|
+
help="A judge the graph asks with 'using NAME', such as "
|
|
395
|
+
"large=ruling:qwen-35b@http://127.0.0.1:8012; repeat for each",
|
|
396
|
+
)
|
|
397
|
+
command.add_argument(
|
|
398
|
+
"--extractor-model",
|
|
399
|
+
help="The generative model that fills a graph's 'extract' fields, as DSPy names it, "
|
|
400
|
+
"such as openai/qwen3.8-27b",
|
|
401
|
+
)
|
|
402
|
+
command.add_argument(
|
|
403
|
+
"--extractor-base-url", help="An OpenAI-compatible endpoint serving --extractor-model"
|
|
404
|
+
)
|
|
405
|
+
|
|
406
|
+
|
|
407
|
+
def judge_spec(text: str) -> tuple[str, str | None, str | None]:
|
|
408
|
+
"""'ruling:qwen-35b@http://127.0.0.1:8012' as its backend, model and endpoint; the model
|
|
409
|
+
and endpoint default to the backend's own, as they do for the graph's judge."""
|
|
410
|
+
rest, at, url = text.partition("@")
|
|
411
|
+
backend, _, model = rest.partition(":")
|
|
412
|
+
if backend not in BACKENDS or (at and not url):
|
|
413
|
+
raise argparse.ArgumentTypeError(
|
|
414
|
+
f"{text!r} is not BACKEND[:MODEL][@URL], with BACKEND one of {sorted(BACKENDS)}"
|
|
415
|
+
)
|
|
416
|
+
return backend, model or None, url or None
|
|
417
|
+
|
|
418
|
+
|
|
419
|
+
def named_judge(text: str) -> tuple[str, str, str | None, str | None]:
|
|
420
|
+
"""'large=ruling:qwen-35b@http://127.0.0.1:8012' as its name, and the judge it names."""
|
|
421
|
+
name, equals, rest = text.partition("=")
|
|
422
|
+
if not equals or not NAME.fullmatch(name):
|
|
423
|
+
raise argparse.ArgumentTypeError(f"{text!r} is not NAME=BACKEND[:MODEL][@URL]")
|
|
424
|
+
return name, *judge_spec(rest)
|
|
425
|
+
|
|
426
|
+
|
|
427
|
+
def cost_arguments(command: argparse.ArgumentParser) -> None:
|
|
428
|
+
"""What each model call costs, for search to prefer the cheaper of two graphs and for
|
|
429
|
+
calibrate to place an escalation line where escalating pays."""
|
|
430
|
+
command.add_argument(
|
|
431
|
+
"--call-cost",
|
|
432
|
+
type=float,
|
|
433
|
+
default=Prices().call,
|
|
434
|
+
help="Score off an example per judge question, to prefer graphs that ask less",
|
|
435
|
+
)
|
|
436
|
+
command.add_argument(
|
|
437
|
+
"--judge-cost",
|
|
438
|
+
type=named_cost,
|
|
439
|
+
action="append",
|
|
440
|
+
default=[],
|
|
441
|
+
metavar="NAME=COST",
|
|
442
|
+
help="Score off an example per question to a named judge; a named judge without one "
|
|
443
|
+
"costs --call-cost",
|
|
444
|
+
)
|
|
445
|
+
command.add_argument(
|
|
446
|
+
"--extract-cost",
|
|
447
|
+
type=float,
|
|
448
|
+
default=Prices().extract,
|
|
449
|
+
help="Score off an example per extraction it runs, to keep only those that pay",
|
|
450
|
+
)
|
|
451
|
+
|
|
452
|
+
|
|
453
|
+
def prices(cli: argparse.ArgumentParser, args: argparse.Namespace) -> Prices:
|
|
454
|
+
named = {name for name, *_ in args.judge}
|
|
455
|
+
for name, _ in args.judge_cost:
|
|
456
|
+
if name not in named:
|
|
457
|
+
cli.error(f"--judge-cost {name} prices a judge no --judge names")
|
|
458
|
+
return Prices(call=args.call_cost, judges=dict(args.judge_cost), extract=args.extract_cost)
|
|
459
|
+
|
|
460
|
+
|
|
461
|
+
def named_cost(text: str) -> tuple[str, float]:
|
|
462
|
+
name, equals, cost = text.partition("=")
|
|
463
|
+
if not equals or not NAME.fullmatch(name):
|
|
464
|
+
raise argparse.ArgumentTypeError(f"{text!r} is not NAME=COST, such as large=0.1")
|
|
465
|
+
try:
|
|
466
|
+
return name, float(cost)
|
|
467
|
+
except ValueError as error:
|
|
468
|
+
raise argparse.ArgumentTypeError(f"the cost in {text!r} is not a number") from error
|
|
469
|
+
|
|
470
|
+
|
|
471
|
+
def graph_models(cli: argparse.ArgumentParser, args, stack: ExitStack, *plans) -> Models:
|
|
472
|
+
"""The models the graphs ask beside their own judge: a client for every judge named on the
|
|
473
|
+
command line, once each name the graphs use is there, and the extractor, if one is given."""
|
|
474
|
+
given = {}
|
|
475
|
+
for name, backend, model, url in args.judge:
|
|
476
|
+
if name in given:
|
|
477
|
+
cli.error(f"--judge {name} is given twice")
|
|
478
|
+
given[name] = backend, model, url
|
|
479
|
+
for name in sorted(set().union(*(plan.judges() for plan in plans if plan)) - given.keys()):
|
|
480
|
+
cli.error(
|
|
481
|
+
f"the graph asks judge {name}; "
|
|
482
|
+
f"say which model with --judge {name}=BACKEND[:MODEL][@URL]"
|
|
483
|
+
)
|
|
484
|
+
judges = {}
|
|
485
|
+
for name, (backend, model, url) in given.items():
|
|
486
|
+
resolved, connection = resolve_backend(backend, model, url)
|
|
487
|
+
retry = RetryPolicy(max_retries=args.retries)
|
|
488
|
+
client = stack.enter_context(connect(backend, connection, args.timeout, retry))
|
|
489
|
+
judges[name] = Judge(client, resolved)
|
|
490
|
+
if args.extractor_base_url and not args.extractor_model:
|
|
491
|
+
cli.error("--extractor-base-url needs an --extractor-model to send extractions to")
|
|
492
|
+
if not args.extractor_model:
|
|
493
|
+
if any(plan.extracts() for plan in plans if plan):
|
|
494
|
+
cli.error(
|
|
495
|
+
"the graph extracts fields; name the model that reads them with --extractor-model"
|
|
496
|
+
)
|
|
497
|
+
return Models(judges)
|
|
498
|
+
options = {"api_base": args.extractor_base_url} if args.extractor_base_url else {}
|
|
499
|
+
lm = dspy.LM(args.extractor_model, temperature=0, cache=False, num_retries=0, **options)
|
|
500
|
+
return Models(judges, Extractor(lm))
|
|
501
|
+
|
|
502
|
+
|
|
503
|
+
def positive(text: str) -> int:
|
|
504
|
+
value = int(text)
|
|
505
|
+
if value < 1:
|
|
506
|
+
raise argparse.ArgumentTypeError(f"must be at least 1, not {value}")
|
|
507
|
+
return value
|
|
508
|
+
|
|
509
|
+
|
|
510
|
+
def judge(args: argparse.Namespace, needed: bool):
|
|
511
|
+
"""The judge model, and a client for it when some graph calls it."""
|
|
512
|
+
model, connection = resolve_backend(args.backend, args.model, args.base_url)
|
|
513
|
+
if not needed:
|
|
514
|
+
return model, nullcontext()
|
|
515
|
+
retry = RetryPolicy(max_retries=args.retries)
|
|
516
|
+
return model, connect(args.backend, connection, args.timeout, retry)
|
|
517
|
+
|
|
518
|
+
|
|
519
|
+
def ablation(cli: argparse.ArgumentParser, args: argparse.Namespace) -> None:
|
|
520
|
+
"""What each input is worth, over the states the task itself supplies."""
|
|
521
|
+
task = load_task(cli, args.task)
|
|
522
|
+
if args.output and args.output.exists():
|
|
523
|
+
cli.error(f"output already exists: {args.output}")
|
|
524
|
+
# A plan is structurally valid once it loads; size limits bound what search may write, and
|
|
525
|
+
# have nothing to say about measuring a graph that already exists.
|
|
526
|
+
plan = read_plan(cli, args.plan, task)
|
|
527
|
+
examples = getattr(task, args.split)
|
|
528
|
+
if not all(hasattr(example, "state") for example in examples):
|
|
529
|
+
cli.error(
|
|
530
|
+
"ablate reads states from a task's examples; a Doom or chess task keeps seeds, "
|
|
531
|
+
"so measure those graphs with rulesmith.ablate.effects over recorded states"
|
|
532
|
+
)
|
|
533
|
+
states = [example.state for example in examples]
|
|
534
|
+
calls = plan.asks_own_judge()
|
|
535
|
+
model, judge_client = judge(args, calls)
|
|
536
|
+
with ExitStack() as stack:
|
|
537
|
+
client = stack.enter_context(judge_client)
|
|
538
|
+
models = graph_models(cli, args, stack, plan)
|
|
539
|
+
execution = ExecutionConfig(max_workers=args.max_workers)
|
|
540
|
+
found = effects(plan, states, client if calls else None, model, execution, models)
|
|
541
|
+
report = {
|
|
542
|
+
"states": len(states),
|
|
543
|
+
"effects": [
|
|
544
|
+
{
|
|
545
|
+
"node": effect.node,
|
|
546
|
+
"kind": effect.kind,
|
|
547
|
+
"changed": effect.changed,
|
|
548
|
+
"share": round(effect.share, 4),
|
|
549
|
+
"reason": effect.reason,
|
|
550
|
+
}
|
|
551
|
+
for effect in found
|
|
552
|
+
],
|
|
553
|
+
}
|
|
554
|
+
if args.output:
|
|
555
|
+
args.output.parent.mkdir(parents=True, exist_ok=True)
|
|
556
|
+
args.output.write_text(dump(report))
|
|
557
|
+
print(dump(report), end="")
|
|
558
|
+
|
|
559
|
+
|
|
560
|
+
def arena(cli: argparse.ArgumentParser, args: argparse.Namespace) -> None:
|
|
561
|
+
task = load_task(cli, args.task)
|
|
562
|
+
if not plays(task, "doom") or not task.deathmatch:
|
|
563
|
+
cli.error("arena requires a Doom task with deathmatch set")
|
|
564
|
+
from rulesmith.arena import Entrant, tournament
|
|
565
|
+
|
|
566
|
+
if args.output.exists():
|
|
567
|
+
cli.error(f"output already exists: {args.output}")
|
|
568
|
+
plans = {}
|
|
569
|
+
for path in args.plans:
|
|
570
|
+
if path.stem in plans:
|
|
571
|
+
cli.error(f"entrants share the name {path.stem}; rename a plan file")
|
|
572
|
+
plans[path.stem] = read_plan(cli, path, task)
|
|
573
|
+
plans[path.stem].check(task.labels, args.max_nodes, args.max_depth)
|
|
574
|
+
if len(plans) < 2:
|
|
575
|
+
cli.error("the arena needs at least two entrants")
|
|
576
|
+
model, judge_client = judge(args, any(plan.asks_own_judge() for plan in plans.values()))
|
|
577
|
+
execution = ExecutionConfig(max_workers=args.max_workers)
|
|
578
|
+
with ExitStack() as stack:
|
|
579
|
+
client = stack.enter_context(judge_client)
|
|
580
|
+
models = graph_models(cli, args, stack, *plans.values())
|
|
581
|
+
entrants = [
|
|
582
|
+
Entrant(name, DecisionProgram(plan, client, model, execution, models=models))
|
|
583
|
+
for name, plan in plans.items()
|
|
584
|
+
]
|
|
585
|
+
report = tournament(task, entrants, getattr(task, args.split))
|
|
586
|
+
report["split"] = args.split
|
|
587
|
+
report["backend"] = args.backend
|
|
588
|
+
args.output.parent.mkdir(parents=True, exist_ok=True)
|
|
589
|
+
args.output.write_text(dump(report))
|
|
590
|
+
print(dump({"ratings": report["ratings"], "output": str(args.output)}), end="")
|
|
591
|
+
|
|
592
|
+
|
|
593
|
+
def decision(cli: argparse.ArgumentParser, path: Path) -> tuple[str, Plan]:
|
|
594
|
+
"""A decision folder's name and graph, read against its task's labels when it has a task, or
|
|
595
|
+
a graph file's stem and graph: the graph as `run` and `serve` would run it."""
|
|
596
|
+
if not path.is_dir():
|
|
597
|
+
return path.stem, read_plan(cli, path)
|
|
598
|
+
file = next((path / f for f in ("graph.rules", "map.json") if (path / f).exists()), None)
|
|
599
|
+
if file is None:
|
|
600
|
+
cli.error(f"{path} holds no graph.rules or map.json")
|
|
601
|
+
task = load_task(cli, path / "task.json") if (path / "task.json").exists() else None
|
|
602
|
+
return path.name, read_plan(cli, file, task)
|
|
603
|
+
|
|
604
|
+
|
|
605
|
+
def watch_drift(cli: argparse.ArgumentParser, args: argparse.Namespace) -> int:
|
|
606
|
+
"""How the decisions a graph made in production compare with what its rules were counted
|
|
607
|
+
on. Only decisions by this exact graph count: a rule's name is its place in the list, so an
|
|
608
|
+
earlier version's rule_2 is another rule."""
|
|
609
|
+
_, plan = decision(cli, args.decision)
|
|
610
|
+
made = digest(plan.model_dump_json())
|
|
611
|
+
files = [
|
|
612
|
+
f
|
|
613
|
+
for path in args.traces
|
|
614
|
+
for f in (sorted(path.rglob("*.jsonl")) if path.is_dir() else [path])
|
|
615
|
+
]
|
|
616
|
+
records = [
|
|
617
|
+
record
|
|
618
|
+
for file in files
|
|
619
|
+
for line in file.read_text().splitlines()
|
|
620
|
+
if line.strip() and (record := json.loads(line))["graph_sha256"] == made
|
|
621
|
+
]
|
|
622
|
+
if not records:
|
|
623
|
+
cli.error(f"no traced decision was made by this graph ({made[:12]})")
|
|
624
|
+
called = Counter(paid for record in records for paid in record.get("called", {}).values())
|
|
625
|
+
report = {
|
|
626
|
+
"graph_sha256": made,
|
|
627
|
+
"decisions": len(records),
|
|
628
|
+
"undecided": sum(record["undecided"] is not None for record in records),
|
|
629
|
+
"calls": dict(sorted(called.items())),
|
|
630
|
+
"drift": drift(plan, records, args.significance),
|
|
631
|
+
}
|
|
632
|
+
print(dump(report), end="")
|
|
633
|
+
return 1 if report["drift"] else 0
|
|
634
|
+
|
|
635
|
+
|
|
636
|
+
def serve_decisions(cli: argparse.ArgumentParser, args: argparse.Namespace) -> None:
|
|
637
|
+
"""Serve each decision under its folder's name, or its file's, until stopped. A folder with a
|
|
638
|
+
task reads its rules against the task's labels, as calibrate and evaluate do."""
|
|
639
|
+
graphs = {}
|
|
640
|
+
for path in args.decisions:
|
|
641
|
+
name, plan = decision(cli, path)
|
|
642
|
+
if name in graphs:
|
|
643
|
+
cli.error(f"two decisions are named {name}; serve them from folders of their own")
|
|
644
|
+
graphs[name] = plan
|
|
645
|
+
model, judge_client = judge(args, any(plan.asks_own_judge() for plan in graphs.values()))
|
|
646
|
+
execution = ExecutionConfig(max_workers=args.max_workers)
|
|
647
|
+
with ExitStack() as stack:
|
|
648
|
+
client = stack.enter_context(judge_client)
|
|
649
|
+
models = graph_models(cli, args, stack, *graphs.values())
|
|
650
|
+
programs = {
|
|
651
|
+
name: DecisionProgram(plan, client, model, execution, args.trace, models=models)
|
|
652
|
+
for name, plan in graphs.items()
|
|
653
|
+
}
|
|
654
|
+
server = stack.enter_context(make_server(programs, args.host, args.port))
|
|
655
|
+
# One line, flushed, so whatever started the server knows when it may call it.
|
|
656
|
+
url = f"http://{args.host}:{server.server_port}"
|
|
657
|
+
print(json.dumps({"url": url, "decisions": sorted(programs)}), flush=True)
|
|
658
|
+
server.serve_forever()
|
|
659
|
+
|
|
660
|
+
|
|
661
|
+
def serve_clef(cli: argparse.ArgumentParser, args: argparse.Namespace) -> None:
|
|
662
|
+
"""Load Clef, then answer its requests until stopped."""
|
|
663
|
+
try:
|
|
664
|
+
# Loading prints whatever the libraries print; stdout's first line must be the URL.
|
|
665
|
+
with redirect_stdout(sys.stderr):
|
|
666
|
+
answer, device = clef.load(args.weights, args.device)
|
|
667
|
+
except ImportError as error:
|
|
668
|
+
cli.error(f"clef-serve needs the clef extra: pip install 'rulesmith[clef]' ({error})")
|
|
669
|
+
with clef.make_server(answer, args.host, args.port) as server:
|
|
670
|
+
# One line, flushed, once the model is loaded, so a caller knows when it may ask.
|
|
671
|
+
url = f"http://{args.host}:{server.server_port}"
|
|
672
|
+
print(json.dumps({"url": url, "weights": args.weights, "device": device}), flush=True)
|
|
673
|
+
server.serve_forever()
|
|
674
|
+
|
|
675
|
+
|
|
676
|
+
def dump(value: dict) -> str:
|
|
677
|
+
return json.dumps(value, indent=2, ensure_ascii=False, allow_nan=False) + "\n"
|
|
678
|
+
|
|
679
|
+
|
|
680
|
+
def draft_labels(cli: argparse.ArgumentParser, args: argparse.Namespace) -> None:
|
|
681
|
+
if args.output.exists():
|
|
682
|
+
cli.error(f"output already exists: {args.output}")
|
|
683
|
+
unlabeled = Unlabeled.model_validate_json(args.inputs.read_text())
|
|
684
|
+
lm_options = {"api_base": args.labeler_base_url} if args.labeler_base_url else {}
|
|
685
|
+
labeler = Labeler(
|
|
686
|
+
dspy.LM(args.labeler_model, cache=False, num_retries=0, **lm_options),
|
|
687
|
+
unlabeled.description,
|
|
688
|
+
unlabeled.labels,
|
|
689
|
+
)
|
|
690
|
+
task = draft_task(
|
|
691
|
+
unlabeled.description,
|
|
692
|
+
unlabeled.labels,
|
|
693
|
+
unlabeled.inputs,
|
|
694
|
+
labeler,
|
|
695
|
+
args.validation,
|
|
696
|
+
args.test,
|
|
697
|
+
args.seed,
|
|
698
|
+
)
|
|
699
|
+
args.output.parent.mkdir(parents=True, exist_ok=True)
|
|
700
|
+
args.output.write_text(task.model_dump_json(indent=2) + "\n")
|
|
701
|
+
counts = Counter(e.label for e in task.examples())
|
|
702
|
+
print(dump({"labels": dict(counts), "output": str(args.output)}), end="")
|
|
703
|
+
|
|
704
|
+
|
|
705
|
+
def save_video(frames: list, fps: int, path: Path) -> None:
|
|
706
|
+
if path.suffix == ".mp4":
|
|
707
|
+
import imageio_ffmpeg
|
|
708
|
+
|
|
709
|
+
height, width, _ = frames[0].shape
|
|
710
|
+
writer = imageio_ffmpeg.write_frames(str(path), (width, height), fps=fps)
|
|
711
|
+
writer.send(None)
|
|
712
|
+
for frame in frames:
|
|
713
|
+
writer.send(frame)
|
|
714
|
+
writer.close()
|
|
715
|
+
return
|
|
716
|
+
from PIL import Image
|
|
717
|
+
|
|
718
|
+
first, *rest = (Image.fromarray(frame) for frame in frames)
|
|
719
|
+
# GIF delays are whole centiseconds, so round to 10 ms rather than truncate.
|
|
720
|
+
delay = round(100 / fps) * 10
|
|
721
|
+
first.save(path, save_all=True, append_images=rest, duration=delay, loop=0)
|
|
722
|
+
|
|
723
|
+
|
|
724
|
+
def scorer(args: argparse.Namespace, task):
|
|
725
|
+
"""What scores a labeled task's answers: its labels, or a grader when one is named. A game
|
|
726
|
+
task's scorer is its games, chosen where the games are set up."""
|
|
727
|
+
if not isinstance(task, Task):
|
|
728
|
+
return None
|
|
729
|
+
if not args.grader_model:
|
|
730
|
+
return Labels(task)
|
|
731
|
+
lm_options = {"api_base": args.grader_base_url} if args.grader_base_url else {}
|
|
732
|
+
lm = dspy.LM(args.grader_model, cache=False, num_retries=0, **lm_options)
|
|
733
|
+
return Graded(Grader(lm, task.description))
|
|
734
|
+
|
|
735
|
+
|
|
736
|
+
def scored_by(args: argparse.Namespace, task) -> str:
|
|
737
|
+
"""How a report's scores were made, since a grader's score is not accuracy."""
|
|
738
|
+
if getattr(args, "grader_model", None):
|
|
739
|
+
return f"grader {args.grader_model}"
|
|
740
|
+
return "labels" if isinstance(task, Task) else "games"
|
|
741
|
+
|
|
742
|
+
|
|
743
|
+
def opponent_of(args: argparse.Namespace) -> str | None:
|
|
744
|
+
"""Who a played task's scores were earned against, since a score means nothing without it."""
|
|
745
|
+
if getattr(args, "engine", None):
|
|
746
|
+
return f"engine {args.engine} at skill {args.engine_skill or 0}"
|
|
747
|
+
return str(args.opponent) if getattr(args, "opponent", None) else None
|
|
748
|
+
|
|
749
|
+
|
|
750
|
+
def write_unit_task(cli: argparse.ArgumentParser, args: argparse.Namespace) -> None:
|
|
751
|
+
"""The task a unit is calibrated and searched on, made by running the whole map, on its
|
|
752
|
+
judge, over the map's own labeled examples."""
|
|
753
|
+
task = load_task(cli, args.task)
|
|
754
|
+
if not isinstance(task, Task):
|
|
755
|
+
cli.error("unit-task reads a labeled task for the whole map")
|
|
756
|
+
if args.output.exists():
|
|
757
|
+
cli.error(f"output already exists: {args.output}")
|
|
758
|
+
plan = read_plan(cli, args.plan, task)
|
|
759
|
+
if not isinstance(plan.nodes.get(args.unit), Graph):
|
|
760
|
+
cli.error(f"{args.plan} has no unit named {args.unit!r}; a unit is a graph node")
|
|
761
|
+
model, judge_client = judge(args, plan.asks_own_judge())
|
|
762
|
+
with ExitStack() as stack:
|
|
763
|
+
client = stack.enter_context(judge_client)
|
|
764
|
+
models = graph_models(cli, args, stack, plan)
|
|
765
|
+
program = DecisionProgram(
|
|
766
|
+
plan, client, model, ExecutionConfig(max_workers=args.max_workers), models=models
|
|
767
|
+
)
|
|
768
|
+
try:
|
|
769
|
+
found = unit_task(program, task, args.unit)
|
|
770
|
+
except MapError as error:
|
|
771
|
+
cli.error(str(error))
|
|
772
|
+
args.output.parent.mkdir(parents=True, exist_ok=True)
|
|
773
|
+
args.output.write_text(found.model_dump_json(indent=2, exclude_none=True) + "\n")
|
|
774
|
+
print(
|
|
775
|
+
dump(
|
|
776
|
+
{
|
|
777
|
+
split: len(getattr(found, split))
|
|
778
|
+
for split in ("train", "validation", "calibration", "test")
|
|
779
|
+
}
|
|
780
|
+
),
|
|
781
|
+
end="",
|
|
782
|
+
)
|
|
783
|
+
|
|
784
|
+
|
|
785
|
+
def calibration(cli: argparse.ArgumentParser, args: argparse.Namespace) -> None:
|
|
786
|
+
"""A plan whose choice questions read as sure as they are right, whose rules draw their lines
|
|
787
|
+
on judgments where the labeled examples put them, and whose rules are counted."""
|
|
788
|
+
task = load_task(cli, args.task)
|
|
789
|
+
if not isinstance(task, Task):
|
|
790
|
+
cli.error("calibrate fits judgments against a labeled task's examples")
|
|
791
|
+
if units(args.plan):
|
|
792
|
+
# Each unit's rules are counted on the cases that reach it, which a map's task does not
|
|
793
|
+
# say; a unit's own task does.
|
|
794
|
+
cli.error(
|
|
795
|
+
"calibrate one unit of a map at a time: write its task with rulesmith unit-task, "
|
|
796
|
+
"then calibrate the unit's own file on it"
|
|
797
|
+
)
|
|
798
|
+
if args.output.exists():
|
|
799
|
+
cli.error(f"output already exists: {args.output}")
|
|
800
|
+
plan = read_plan(cli, args.plan, task)
|
|
801
|
+
priced = prices(cli, args)
|
|
802
|
+
model, judge_client = judge(args, plan.asks_own_judge())
|
|
803
|
+
held_out = [e for e in task.validation if e.label is not None]
|
|
804
|
+
with ExitStack() as stack:
|
|
805
|
+
client = stack.enter_context(judge_client)
|
|
806
|
+
# Every pass asks the same questions; the later ones reread the first's answers.
|
|
807
|
+
remembered = Remembered(client)
|
|
808
|
+
models = graph_models(cli, args, stack, plan).remembered()
|
|
809
|
+
fitted, temperatures = calibrate(plan, task, remembered, model, models)
|
|
810
|
+
before = correct(fitted, held_out, remembered, model, models)
|
|
811
|
+
fitted, cutoffs = fit_cutoffs(fitted, task, remembered, model, models, priced)
|
|
812
|
+
after = correct(fitted, held_out, remembered, model, models)
|
|
813
|
+
fitted = measure(fitted, task, remembered, model, models)
|
|
814
|
+
written = (
|
|
815
|
+
render_rules(fitted, task.criteria())
|
|
816
|
+
if args.output.suffix == ".rules"
|
|
817
|
+
else fitted.model_dump_json(indent=2)
|
|
818
|
+
)
|
|
819
|
+
if args.output.suffix == ".rules" and args.plan.suffix == ".rules":
|
|
820
|
+
written = keep_comments(args.plan.read_text(), written)
|
|
821
|
+
args.output.parent.mkdir(parents=True, exist_ok=True)
|
|
822
|
+
args.output.write_text(written + "\n")
|
|
823
|
+
counted = {
|
|
824
|
+
name: node.measured.model_dump()
|
|
825
|
+
for name, node in fitted.nodes.items()
|
|
826
|
+
if getattr(node, "measured", None) is not None
|
|
827
|
+
}
|
|
828
|
+
# Fitted on examples search never saw, or on the ones it chose the graph by.
|
|
829
|
+
report = {
|
|
830
|
+
"fitted_on": task.fitted_on(),
|
|
831
|
+
"temperatures": temperatures,
|
|
832
|
+
"cutoffs": cutoffs,
|
|
833
|
+
"rules": counted,
|
|
834
|
+
}
|
|
835
|
+
if priced != Prices():
|
|
836
|
+
# A line placed against prices holds for those prices; say which they were.
|
|
837
|
+
report["prices"] = priced.model_dump()
|
|
838
|
+
if held_out:
|
|
839
|
+
# Validation is scored, never fitted on, so this is what the fit is worth on fresh examples.
|
|
840
|
+
report["validation"] = {
|
|
841
|
+
"before": sum(before) / len(held_out),
|
|
842
|
+
"after": sum(after) / len(held_out),
|
|
843
|
+
}
|
|
844
|
+
print(dump(report | {"output": str(args.output)}), end="")
|
|
845
|
+
|
|
846
|
+
|
|
847
|
+
def read_plan(
|
|
848
|
+
cli: argparse.ArgumentParser, path: Path, task: "Task | DoomTask | ChessTask | None" = None
|
|
849
|
+
) -> Plan:
|
|
850
|
+
"""A graph from its JSON, or from the rules it is written as, which may ask over the labels
|
|
851
|
+
of the task the command names. A file that does not hold one is a mistake in the command,
|
|
852
|
+
reported as such rather than as a crash."""
|
|
853
|
+
try:
|
|
854
|
+
return load(path, described_labels(task) if task else None)
|
|
855
|
+
except (ValueError, OSError) as error:
|
|
856
|
+
cli.error(f"{path}: {error}")
|
|
857
|
+
|
|
858
|
+
|
|
859
|
+
def game(cli: argparse.ArgumentParser, name: str) -> ModuleType:
|
|
860
|
+
"""A game's module. Its packages come with the extra of the same name, which a core install
|
|
861
|
+
leaves out, so a missing one is a matter of installing rather than a crash."""
|
|
862
|
+
try:
|
|
863
|
+
return importlib.import_module(f"rulesmith.{name}")
|
|
864
|
+
except ModuleNotFoundError as error:
|
|
865
|
+
cli.error(f"a {name} task needs {error.name}: pip install 'rulesmith[{name}]'")
|
|
866
|
+
|
|
867
|
+
|
|
868
|
+
def plays(task, name: str) -> bool:
|
|
869
|
+
"""Whether a task is one of a game's, asked without importing a game that may be missing."""
|
|
870
|
+
return type(task).__module__ == f"rulesmith.{name}"
|
|
871
|
+
|
|
872
|
+
|
|
873
|
+
def load_task(cli: argparse.ArgumentParser, path: Path) -> "Task | DoomTask | ChessTask":
|
|
874
|
+
try:
|
|
875
|
+
data = json.loads(path.read_text())
|
|
876
|
+
if data.get("game") == "chess":
|
|
877
|
+
kind = game(cli, "chess").ChessTask
|
|
878
|
+
elif "scenario" in data:
|
|
879
|
+
kind = game(cli, "doom").DoomTask
|
|
880
|
+
else:
|
|
881
|
+
kind = Task
|
|
882
|
+
return kind.model_validate(data)
|
|
883
|
+
except ValueError as error:
|
|
884
|
+
cli.error(f"{path}: {error}")
|
|
885
|
+
|
|
886
|
+
|
|
887
|
+
def read_state(cli: argparse.ArgumentParser, path: Path):
|
|
888
|
+
try:
|
|
889
|
+
return json.loads(path.read_text())
|
|
890
|
+
except (OSError, json.JSONDecodeError) as error:
|
|
891
|
+
cli.error(f"{path}: {error}")
|
|
892
|
+
|
|
893
|
+
|
|
894
|
+
def main() -> None:
|
|
895
|
+
# The default SIGTERM action skips cleanup, which leaves ViZDoom engines running (they
|
|
896
|
+
# ignore SIGTERM themselves); exiting normally closes every environment on the way out.
|
|
897
|
+
signal.signal(signal.SIGTERM, lambda number, _: sys.exit(128 + number))
|
|
898
|
+
cli = parser()
|
|
899
|
+
argv = sys.argv[1:]
|
|
900
|
+
# Arguments after -- belong to `ruling train`. Older Python 3.12 releases match an empty
|
|
901
|
+
# positional before the options and then refuse them, so they are split off here instead.
|
|
902
|
+
split = argv.index("--") if "--" in argv else len(argv)
|
|
903
|
+
args = cli.parse_args(argv[:split])
|
|
904
|
+
args.train_args = argv[split + 1 :]
|
|
905
|
+
if args.train_args and args.command not in ("tune", "distill"):
|
|
906
|
+
cli.error("only tune and distill pass arguments after -- on to ruling train")
|
|
907
|
+
try:
|
|
908
|
+
command(cli, args)
|
|
909
|
+
except (TypeSafeError, ChatJudgeError) as error:
|
|
910
|
+
# A missing key or an unreachable judge is something to set or start, said in a line.
|
|
911
|
+
cli.error(str(error))
|
|
912
|
+
|
|
913
|
+
|
|
914
|
+
def command(cli: argparse.ArgumentParser, args: argparse.Namespace) -> None:
|
|
915
|
+
if args.command == "draft-labels":
|
|
916
|
+
draft_labels(cli, args)
|
|
917
|
+
return
|
|
918
|
+
if args.command == "draw":
|
|
919
|
+
print(mermaid(read_plan(cli, args.plan), closed=args.closed))
|
|
920
|
+
return
|
|
921
|
+
if args.command in ("tune", "distill"):
|
|
922
|
+
task = load_task(cli, args.task)
|
|
923
|
+
if not isinstance(task, Task):
|
|
924
|
+
cli.error(f"{args.command} needs a task's inputs; a game has none to learn from")
|
|
925
|
+
extra = list(args.train_args)
|
|
926
|
+
try:
|
|
927
|
+
if args.command == "distill":
|
|
928
|
+
# Refused before the teacher is paid, and the teacher asked before anything is
|
|
929
|
+
# written, so a teacher that cannot answer leaves no half-made tuning set.
|
|
930
|
+
taken = [args.output / f"{name}.jsonl" for name in FILES]
|
|
931
|
+
if any(path.exists() for path in taken):
|
|
932
|
+
cli.error(f"already exists: {args.output}; distill into a new directory")
|
|
933
|
+
backend, model, url = args.teacher
|
|
934
|
+
teacher, connection = resolve_backend(backend, model, url)
|
|
935
|
+
with connect(backend, connection, args.timeout, RetryPolicy()) as client:
|
|
936
|
+
found = taught(task, client, teacher)
|
|
937
|
+
report = tune(task, args.output, args.ruling_dir, args.base, extra, found)
|
|
938
|
+
report["teacher"] = teacher
|
|
939
|
+
else:
|
|
940
|
+
report = tune(task, args.output, args.ruling_dir, args.base, extra)
|
|
941
|
+
except (FileExistsError, FileNotFoundError) as error:
|
|
942
|
+
cli.error(str(error))
|
|
943
|
+
except subprocess.CalledProcessError as error:
|
|
944
|
+
cli.error(f"ruling train failed with exit code {error.returncode}")
|
|
945
|
+
print(dump(report), end="")
|
|
946
|
+
return
|
|
947
|
+
if args.command == "arena":
|
|
948
|
+
arena(cli, args)
|
|
949
|
+
return
|
|
950
|
+
if args.command == "ablate":
|
|
951
|
+
ablation(cli, args)
|
|
952
|
+
return
|
|
953
|
+
if args.command == "unit-task":
|
|
954
|
+
return write_unit_task(cli, args)
|
|
955
|
+
if args.command == "calibrate":
|
|
956
|
+
calibration(cli, args)
|
|
957
|
+
return
|
|
958
|
+
if args.command == "serve":
|
|
959
|
+
serve_decisions(cli, args)
|
|
960
|
+
return
|
|
961
|
+
if args.command == "clef-serve":
|
|
962
|
+
serve_clef(cli, args)
|
|
963
|
+
return
|
|
964
|
+
if args.command == "drift":
|
|
965
|
+
sys.exit(watch_drift(cli, args))
|
|
966
|
+
task = None if args.command == "run" else load_task(cli, args.task)
|
|
967
|
+
plan = read_plan(cli, args.plan, task) if args.plan else task.baseline_plan()
|
|
968
|
+
whole, unit = None, getattr(args, "unit", None)
|
|
969
|
+
map_task = getattr(args, "map_task", None)
|
|
970
|
+
if map_task is not None and unit is None:
|
|
971
|
+
cli.error("--map-task scores the map around a --unit; name the unit")
|
|
972
|
+
if map_task is not None:
|
|
973
|
+
map_task = load_task(cli, map_task)
|
|
974
|
+
if not isinstance(map_task, Task):
|
|
975
|
+
cli.error("--map-task must be a labeled task for the whole map")
|
|
976
|
+
if unit is not None:
|
|
977
|
+
# The map is what the winner must fit; the unit's own graph is what search starts from.
|
|
978
|
+
whole = plan
|
|
979
|
+
if not isinstance(whole.nodes.get(unit), Graph):
|
|
980
|
+
cli.error(f"{args.plan} has no unit named {unit!r}; a unit is a graph node")
|
|
981
|
+
plan = whole.nodes[unit].plan
|
|
982
|
+
if args.command != "run":
|
|
983
|
+
plan.check(task.labels, args.max_nodes, args.max_depth)
|
|
984
|
+
if args.command == "validate":
|
|
985
|
+
checked = {"nodes": len(plan.nodes), "labels": task.labels}
|
|
986
|
+
# Valid, since the inputs a task sees may always carry what the rules read, but worth
|
|
987
|
+
# knowing: an input that does not is an error at run time, not an answer.
|
|
988
|
+
if unanswered := plan.unanswered():
|
|
989
|
+
checked["unanswered"] = unanswered
|
|
990
|
+
# A map's units, so what was validated across files is on record.
|
|
991
|
+
if named := units(args.plan):
|
|
992
|
+
checked["units"] = named
|
|
993
|
+
print(dump(checked), end="")
|
|
994
|
+
return
|
|
995
|
+
if getattr(args, "visible", False) and not plays(task, "doom"):
|
|
996
|
+
cli.error("--visible requires a Doom task")
|
|
997
|
+
deathmatch = args.command != "run" and plays(task, "doom") and task.deathmatch
|
|
998
|
+
played = deathmatch or (args.command != "run" and plays(task, "chess"))
|
|
999
|
+
engine = getattr(args, "engine", None)
|
|
1000
|
+
tuned = [flag for flag in ("engine_skill", "engine_seconds") if getattr(args, flag, None)]
|
|
1001
|
+
if tuned and not engine:
|
|
1002
|
+
cli.error("--engine-skill and --engine-seconds set the strength of an --engine")
|
|
1003
|
+
if engine and not plays(task, "chess"):
|
|
1004
|
+
cli.error("--engine plays a chess task; other games are played against a graph")
|
|
1005
|
+
if engine and args.opponent:
|
|
1006
|
+
cli.error("a chess task is played against --opponent or --engine, not both")
|
|
1007
|
+
against = bool(args.opponent or engine) if hasattr(args, "opponent") else False
|
|
1008
|
+
if args.command in ("evaluate", "optimize", "record") and played != against:
|
|
1009
|
+
cli.error("--opponent is required for, and only for, tasks played against another graph")
|
|
1010
|
+
opponent = None
|
|
1011
|
+
if played and not engine:
|
|
1012
|
+
opponent = read_plan(cli, args.opponent, task)
|
|
1013
|
+
opponent.check(task.labels, args.max_nodes, args.max_depth)
|
|
1014
|
+
if getattr(args, "parallel_examples", 1) > 1 and plays(task, "doom") and not deathmatch:
|
|
1015
|
+
cli.error("--parallel-examples needs a deathmatch task; other Doom episodes share one game")
|
|
1016
|
+
if args.command == "record" and not plays(task, "doom"):
|
|
1017
|
+
cli.error("record requires a Doom task")
|
|
1018
|
+
if args.command == "record" and args.output.suffix not in (".gif", ".mp4"):
|
|
1019
|
+
cli.error("record writes .gif or .mp4")
|
|
1020
|
+
grading = getattr(args, "grader_model", None)
|
|
1021
|
+
if getattr(args, "grader_base_url", None) and not grading:
|
|
1022
|
+
cli.error("--grader-base-url needs a --grader-model to send grades to")
|
|
1023
|
+
if grading and not isinstance(task, Task):
|
|
1024
|
+
cli.error("--grader-model grades a labeled task's examples; games score themselves")
|
|
1025
|
+
if grading and model_name(grading) == model_name(resolve_backend(args.backend, args.model)[0]):
|
|
1026
|
+
# Search would learn to ask the grader's own question and score perfectly for nothing.
|
|
1027
|
+
cli.error("--grader-model is the judge the graph asks; grade with a different model")
|
|
1028
|
+
if getattr(args, "visible", False) and deathmatch:
|
|
1029
|
+
cli.error("--visible shows a single-player Doom game; a deathmatch has no one window")
|
|
1030
|
+
keeping = getattr(args, "keep_judge", None) is not None
|
|
1031
|
+
if keeping and (grading or not isinstance(task, Task)):
|
|
1032
|
+
cli.error("--keep-judge measures rules against a labeled task's labels")
|
|
1033
|
+
if args.command in ("evaluate", "optimize") and isinstance(task, Task) and not grading:
|
|
1034
|
+
splits = ("train", "validation") if args.command == "optimize" else (args.split,)
|
|
1035
|
+
unlabeled = [e.id for split in splits for e in getattr(task, split) if e.label is None]
|
|
1036
|
+
if unlabeled:
|
|
1037
|
+
cli.error(f"examples {unlabeled} have no label; label them or pass a --grader-model")
|
|
1038
|
+
if args.command in ("evaluate", "optimize", "record"):
|
|
1039
|
+
resuming = getattr(args, "resume", False)
|
|
1040
|
+
checkpoint = args.output / "search" / "gepa_state.bin"
|
|
1041
|
+
if resuming and not checkpoint.exists():
|
|
1042
|
+
cli.error(f"no checkpoint to resume in {args.output}: {checkpoint} is missing")
|
|
1043
|
+
if args.output.exists() and not resuming:
|
|
1044
|
+
hint = "; pass --resume to continue its search" if checkpoint.exists() else ""
|
|
1045
|
+
cli.error(f"output already exists: {args.output}{hint}")
|
|
1046
|
+
args.output.parent.mkdir(parents=True, exist_ok=True)
|
|
1047
|
+
if args.command == "optimize":
|
|
1048
|
+
priced = prices(cli, args)
|
|
1049
|
+
config = SearchConfig(
|
|
1050
|
+
max_nodes=args.max_nodes,
|
|
1051
|
+
max_depth=args.max_depth,
|
|
1052
|
+
max_metric_calls=args.max_metric_calls,
|
|
1053
|
+
reflection_batch=args.reflection_batch,
|
|
1054
|
+
repair_attempts=args.repair_attempts,
|
|
1055
|
+
seed=args.seed,
|
|
1056
|
+
parallel_examples=args.parallel_examples,
|
|
1057
|
+
merges=args.merges,
|
|
1058
|
+
call_cost=priced.call,
|
|
1059
|
+
call_costs=priced.judges,
|
|
1060
|
+
extract_cost=priced.extract,
|
|
1061
|
+
significance=args.significance,
|
|
1062
|
+
mine=args.mine,
|
|
1063
|
+
evidence=args.evidence,
|
|
1064
|
+
keep_judge=args.keep_judge,
|
|
1065
|
+
)
|
|
1066
|
+
lm_options = {"api_base": args.reflection_base_url} if args.reflection_base_url else {}
|
|
1067
|
+
if args.reflection_timeout is not None:
|
|
1068
|
+
lm_options["timeout"] = args.reflection_timeout
|
|
1069
|
+
if args.reflection_options:
|
|
1070
|
+
lm_options["extra_body"] = args.reflection_options
|
|
1071
|
+
reflection = dspy.LM(
|
|
1072
|
+
args.reflection_model,
|
|
1073
|
+
temperature=1,
|
|
1074
|
+
max_tokens=args.reflection_max_tokens,
|
|
1075
|
+
cache=False,
|
|
1076
|
+
num_retries=0,
|
|
1077
|
+
**lm_options,
|
|
1078
|
+
)
|
|
1079
|
+
|
|
1080
|
+
reflect = GraphProposer(reflection)
|
|
1081
|
+
|
|
1082
|
+
execution = ExecutionConfig(max_workers=args.max_workers)
|
|
1083
|
+
model, judge_client = judge(
|
|
1084
|
+
args,
|
|
1085
|
+
args.command == "optimize"
|
|
1086
|
+
or any(graph.asks_own_judge() for graph in (plan, opponent) if graph),
|
|
1087
|
+
)
|
|
1088
|
+
with ExitStack() as stack:
|
|
1089
|
+
client = stack.enter_context(judge_client)
|
|
1090
|
+
models = graph_models(cli, args, stack, plan, opponent)
|
|
1091
|
+
score = scorer(args, task) if args.command in ("evaluate", "optimize") else None
|
|
1092
|
+
if deathmatch:
|
|
1093
|
+
from rulesmith.arena import Duels
|
|
1094
|
+
|
|
1095
|
+
score = Duels(task, DecisionProgram(opponent, client, model, execution, models=models))
|
|
1096
|
+
elif args.command != "run" and plays(task, "chess"):
|
|
1097
|
+
from rulesmith.chess import Engine, Games
|
|
1098
|
+
|
|
1099
|
+
if engine:
|
|
1100
|
+
rival = Engine(engine, args.engine_skill or 0, args.engine_seconds or 0.05)
|
|
1101
|
+
score = Games(task, stack.enter_context(rival))
|
|
1102
|
+
else:
|
|
1103
|
+
score = Games(
|
|
1104
|
+
task, DecisionProgram(opponent, client, model, execution, models=models)
|
|
1105
|
+
)
|
|
1106
|
+
elif args.command != "run" and plays(task, "doom"):
|
|
1107
|
+
from rulesmith.doom import DoomEnvironment, Episodes
|
|
1108
|
+
|
|
1109
|
+
env = DoomEnvironment(
|
|
1110
|
+
task, visible=getattr(args, "visible", False), record=args.command == "record"
|
|
1111
|
+
)
|
|
1112
|
+
score = Episodes(task, stack.enter_context(env))
|
|
1113
|
+
if args.command == "run":
|
|
1114
|
+
prediction = DecisionProgram(plan, client, model, execution, args.trace, models=models)(
|
|
1115
|
+
state=read_state(cli, args.state)
|
|
1116
|
+
)
|
|
1117
|
+
print(dump(dict(prediction)), end="")
|
|
1118
|
+
elif args.command == "record":
|
|
1119
|
+
from vizdoom import DEFAULT_TICRATE
|
|
1120
|
+
|
|
1121
|
+
from rulesmith.arena import duel
|
|
1122
|
+
|
|
1123
|
+
program = DecisionProgram(plan, client, model, execution, models=models)
|
|
1124
|
+
frames, episodes = [], []
|
|
1125
|
+
for seed in args.seed:
|
|
1126
|
+
if deathmatch:
|
|
1127
|
+
played, _ = duel(task, program, score.opponent, seed, True, frames)
|
|
1128
|
+
episodes.append(played.record)
|
|
1129
|
+
else:
|
|
1130
|
+
episodes.append(score(program, seed).record)
|
|
1131
|
+
frames += env.frames
|
|
1132
|
+
save_video(frames, DEFAULT_TICRATE, args.output)
|
|
1133
|
+
print(dump({"episodes": episodes, "output": str(args.output)}), end="")
|
|
1134
|
+
elif args.command == "evaluate":
|
|
1135
|
+
report = evaluate_plan(
|
|
1136
|
+
plan,
|
|
1137
|
+
getattr(task, args.split),
|
|
1138
|
+
client,
|
|
1139
|
+
model,
|
|
1140
|
+
score,
|
|
1141
|
+
execution,
|
|
1142
|
+
args.parallel_examples,
|
|
1143
|
+
models,
|
|
1144
|
+
)
|
|
1145
|
+
report["split"] = args.split
|
|
1146
|
+
report["backend"] = args.backend
|
|
1147
|
+
report["model"] = model
|
|
1148
|
+
report["scored_by"] = scored_by(args, task)
|
|
1149
|
+
report["opponent"] = opponent_of(args)
|
|
1150
|
+
summary = {"score": report["score"], "output": str(args.output)}
|
|
1151
|
+
if isinstance(task, Task):
|
|
1152
|
+
real = at_prevalence(task, report["examples"])
|
|
1153
|
+
if real is not None:
|
|
1154
|
+
report["at_prevalence"] = summary["at_prevalence"] = real
|
|
1155
|
+
report["drift"] = summary["drift"] = drift(plan, report["examples"], args.drift)
|
|
1156
|
+
for moved in report["drift"]:
|
|
1157
|
+
print(
|
|
1158
|
+
f"warning: {moved['rule']} decided {moved['share']:.0%} of these "
|
|
1159
|
+
f"examples against {moved['measured_share']:.0%} when calibrated "
|
|
1160
|
+
f"(p = {moved['p']:.2g}); the inputs have moved from the ones it was "
|
|
1161
|
+
"measured on",
|
|
1162
|
+
file=sys.stderr,
|
|
1163
|
+
)
|
|
1164
|
+
args.output.write_text(dump(report))
|
|
1165
|
+
print(dump(summary), end="")
|
|
1166
|
+
else:
|
|
1167
|
+
with redirect_stdout(sys.stderr):
|
|
1168
|
+
best, report = optimize_plan(
|
|
1169
|
+
task,
|
|
1170
|
+
plan,
|
|
1171
|
+
client,
|
|
1172
|
+
model,
|
|
1173
|
+
reflect,
|
|
1174
|
+
config,
|
|
1175
|
+
execution,
|
|
1176
|
+
score,
|
|
1177
|
+
run_dir=args.output / "search",
|
|
1178
|
+
models=models,
|
|
1179
|
+
)
|
|
1180
|
+
report["backend"] = args.backend
|
|
1181
|
+
report["reflection_model"] = args.reflection_model
|
|
1182
|
+
report["scored_by"] = scored_by(args, task)
|
|
1183
|
+
report["opponent"] = opponent_of(args)
|
|
1184
|
+
if whole is not None:
|
|
1185
|
+
# Placed where the unit was, the winner must meet every hand-off the map makes.
|
|
1186
|
+
placed = whole.nodes[unit].model_copy(update={"plan": best})
|
|
1187
|
+
try:
|
|
1188
|
+
placed_map = Plan.model_validate(
|
|
1189
|
+
whole.model_copy(
|
|
1190
|
+
update={"nodes": whole.nodes | {unit: placed}}
|
|
1191
|
+
).model_dump()
|
|
1192
|
+
)
|
|
1193
|
+
except ValueError as error:
|
|
1194
|
+
cli.error(f"the graph search found does not fit the map at {unit!r}: {error}")
|
|
1195
|
+
report["unit"] = {"name": unit, "path": units(args.plan).get(unit), "placed": True}
|
|
1196
|
+
if map_task is not None:
|
|
1197
|
+
# A unit is judged by what the whole map does with it, on the map's own
|
|
1198
|
+
# examples: better alone and worse for the map is worse.
|
|
1199
|
+
scores = {}
|
|
1200
|
+
for when, graph in (("before", whole), ("after", placed_map)):
|
|
1201
|
+
scored = evaluate_plan(
|
|
1202
|
+
graph,
|
|
1203
|
+
map_task.validation,
|
|
1204
|
+
client,
|
|
1205
|
+
model,
|
|
1206
|
+
Labels(map_task),
|
|
1207
|
+
execution,
|
|
1208
|
+
args.parallel_examples,
|
|
1209
|
+
models,
|
|
1210
|
+
)
|
|
1211
|
+
scores[when] = scored["score"]
|
|
1212
|
+
if scores["after"] < scores["before"]:
|
|
1213
|
+
cli.error(
|
|
1214
|
+
f"placed in the map, the graph search found scores {scores['after']:g} "
|
|
1215
|
+
f"on the map task's validation examples against {scores['before']:g} "
|
|
1216
|
+
"before: better alone, worse for the map"
|
|
1217
|
+
)
|
|
1218
|
+
report["unit"]["map_validation"] = scores
|
|
1219
|
+
args.output.mkdir(exist_ok=True)
|
|
1220
|
+
(args.output / "plan.json").write_text(best.model_dump_json(indent=2) + "\n")
|
|
1221
|
+
# The same graph as the rules search wrote, which is the form a person edits and
|
|
1222
|
+
# any command accepts; a graph beyond rules has no such form, and the report says.
|
|
1223
|
+
try:
|
|
1224
|
+
written = render_rules(best, described_labels(task))
|
|
1225
|
+
(args.output / "plan.rules").write_text(written + "\n")
|
|
1226
|
+
report["rules"] = "plan.rules"
|
|
1227
|
+
except RuleError:
|
|
1228
|
+
report["rules"] = None
|
|
1229
|
+
(args.output / "report.json").write_text(dump(report))
|
|
1230
|
+
print(dump(report), end="")
|
|
1231
|
+
|
|
1232
|
+
|
|
1233
|
+
if __name__ == "__main__":
|
|
1234
|
+
main()
|