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/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()