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/doom.py ADDED
@@ -0,0 +1,550 @@
1
+ """Score decision graphs by playing ViZDoom episodes, one graph execution per action."""
2
+
3
+ import math
4
+ import os
5
+ import tempfile
6
+ from collections import Counter, deque
7
+ from functools import cache
8
+ from pathlib import Path
9
+ from typing import Literal, Self
10
+
11
+ import numpy as np
12
+ import vizdoom as vzd
13
+ from pydantic import ConfigDict, Field, model_validator
14
+
15
+ from rulesmith.graph import Plan, StrictModel, Text
16
+ from rulesmith.level import Level
17
+ from rulesmith.optimize import Outcome
18
+ from rulesmith.runtime import DecisionProgram, undecided
19
+
20
+
21
+ class Episode(StrictModel):
22
+ """One episode: a seed, optionally on another map or scenario than the task's own."""
23
+
24
+ model_config = ConfigDict(extra="forbid", frozen=True)
25
+ seed: int
26
+ map: Text | None = None
27
+ scenario: Text | None = None
28
+
29
+
30
+ class DoomTask(StrictModel):
31
+ description: Text
32
+ scenario: Text = Field(description="Name of a scenario bundled with ViZDoom, e.g. basic")
33
+ labels: list[Text] = Field(
34
+ min_length=2,
35
+ description="Actions the graph may choose; join buttons with + to press them together",
36
+ )
37
+ reward_range: tuple[float, float] = Field(description="Lowest and highest episode reward")
38
+ tics_per_action: int = Field(default=4, ge=1)
39
+ episode_timeout: int | None = Field(
40
+ default=None, ge=1, description="Tics before an episode ends; defaults to the scenario's"
41
+ )
42
+ feedback_steps: int = Field(default=8, ge=1)
43
+ map: Text | None = Field(default=None, description="Map to load, e.g. E1M1")
44
+ deathmatch: bool = Field(
45
+ default=False,
46
+ description="Played only as networked duels in the arena; deathmatch maps such as "
47
+ "cig's have no single-player start",
48
+ )
49
+ variables: list[Text] = Field(
50
+ default_factory=list, description="Game variables to add, e.g. HEALTH or AMMO2"
51
+ )
52
+ observe: list[Literal["pose", "rangefinder", "compass", "memory", "target", "route"]] = Field(
53
+ default_factory=list, description="Extra observations beyond visible objects"
54
+ )
55
+ rays: int = Field(default=7, ge=2, description="Rangefinder directions across the view")
56
+ rangefinder_range: float = Field(
57
+ default=2048, gt=0, description="Map units reported when no wall is within reach"
58
+ )
59
+ compass_targets: list[Text] = Field(
60
+ default_factory=list, description="Object names the compass tracks anywhere on the map"
61
+ )
62
+ compass_size: int = Field(default=8, ge=1, description="Nearest compass targets reported")
63
+ memory_steps: int = Field(default=4, ge=1, description="Recent actions and movement kept")
64
+ cell_size: float = Field(default=64, gt=0, description="Map units per exploration cell")
65
+ rewards: dict[Text, float] = Field(
66
+ default_factory=dict,
67
+ description="ViZDoom rewards by setter name, e.g. kill_reward, plus explore per new "
68
+ "cell and progress per map unit the route to the exit shrinks",
69
+ )
70
+ train: list[int | Episode] = Field(min_length=1)
71
+ validation: list[int | Episode] = Field(min_length=1)
72
+ test: list[int | Episode] = Field(min_length=1)
73
+
74
+ @model_validator(mode="after")
75
+ def consistent(self) -> Self:
76
+ if len(set(self.labels)) != len(self.labels):
77
+ raise ValueError("labels must be distinct")
78
+ if self.reward_range[0] >= self.reward_range[1]:
79
+ raise ValueError("reward_range must be increasing")
80
+ episodes = [self.episode(e) for e in self.train + self.validation + self.test]
81
+ if len(set(episodes)) != len(episodes):
82
+ raise ValueError("episodes must be distinct across splits")
83
+ missing = {name for name in self.variables if not hasattr(vzd.GameVariable, name)}
84
+ if missing:
85
+ raise ValueError(f"unknown ViZDoom game variables: {sorted(missing)}")
86
+ if "compass" in self.observe and not self.compass_targets:
87
+ raise ValueError("the compass needs compass_targets")
88
+ unknown = {
89
+ name
90
+ for name in self.rewards
91
+ if name not in CUSTOM_REWARDS
92
+ and not (
93
+ name.endswith(("_reward", "_penalty")) and hasattr(vzd.DoomGame, f"set_{name}")
94
+ )
95
+ }
96
+ if unknown:
97
+ raise ValueError(f"unknown ViZDoom rewards: {sorted(unknown)}")
98
+ return self
99
+
100
+ def episode(self, episode: int | Episode) -> Episode:
101
+ """An episode with its map and scenario filled in from the task."""
102
+ if isinstance(episode, int):
103
+ episode = Episode(seed=episode)
104
+ return Episode(
105
+ seed=episode.seed,
106
+ map=episode.map or self.map,
107
+ scenario=episode.scenario or self.scenario,
108
+ )
109
+
110
+ def baseline_plan(self) -> Plan:
111
+ """A call-free seed that always presses the first listed button."""
112
+ return Plan.constant(self.labels[0])
113
+
114
+
115
+ # Doom's horizontal field of view, and the smallest opening the player fits through.
116
+ FIELD_OF_VIEW = 90
117
+ PLAYER_HEIGHT = 56
118
+
119
+
120
+ def cross(a: np.ndarray, b: np.ndarray) -> np.ndarray:
121
+ """The z component of 2-D cross products; NumPy 2 no longer accepts 2-D vectors."""
122
+ return a[..., 0] * b[..., 1] - a[..., 1] * b[..., 0]
123
+
124
+
125
+ def walls(state) -> tuple[np.ndarray, np.ndarray]:
126
+ """Segments that stop the player: solid walls, and gaps too low to pass such as closed doors."""
127
+ spans: dict[tuple, list] = {}
128
+ blocking: dict[tuple, bool] = {}
129
+ for sector in state.sectors:
130
+ for line in sector.lines:
131
+ key = tuple(sorted([(line.x1, line.y1), (line.x2, line.y2)]))
132
+ spans.setdefault(key, []).append(sector.ceiling_height - sector.floor_height)
133
+ blocking[key] = blocking.get(key, False) or line.is_blocking
134
+ solid = [key for key, gaps in spans.items() if blocking[key] or min(gaps) < PLAYER_HEIGHT]
135
+ segments = np.array(solid, dtype=float).reshape(-1, 2, 2)
136
+ return segments[:, 0], segments[:, 1]
137
+
138
+
139
+ CUSTOM_REWARDS = {"explore", "progress"}
140
+ # Deathmatch rules from ViZDoom's competition examples, for whoever hosts a networked game,
141
+ # minus +viz_nocheat: it switches off the labels and objects that graphs observe.
142
+ DEATHMATCH = (
143
+ "-deathmatch -nomonsters +sv_forcerespawn 1 +sv_noautoaim 1 +sv_respawnprotect 1 "
144
+ "+sv_spawnfarthest 1 +sv_nocrouch 1"
145
+ )
146
+
147
+
148
+ def scenario_config(scenario: str) -> str:
149
+ config = os.path.join(vzd.scenarios_path, f"{scenario}.cfg")
150
+ if not os.path.isfile(config):
151
+ raise ValueError(f"unknown ViZDoom scenario: {scenario}")
152
+ return config
153
+
154
+
155
+ def setting(config: str, key: str) -> str | None:
156
+ """A ViZDoom config value; keys ignore case and underscores, as ViZDoom does."""
157
+ wanted = key.replace("_", "").lower()
158
+ for line in Path(config).read_text().splitlines():
159
+ name, _, value = line.split("#", 1)[0].partition("=")
160
+ if value and name.strip().replace("_", "").lower() == wanted:
161
+ return value.strip()
162
+ return None
163
+
164
+
165
+ @cache
166
+ def load_level(path: Path, map_name: str) -> Level:
167
+ """Building a level's route grid takes seconds, and levels never change once built."""
168
+ return Level(path, map_name)
169
+
170
+
171
+ def level_file(config: str, map_name: str) -> Path:
172
+ """The WAD holding a scenario's map: its scenario WAD, else its game WAD."""
173
+ here = Path(config).parent
174
+ for key in ("doom_scenario_path", "doom_game_path"):
175
+ name = setting(config, key)
176
+ if name is None:
177
+ continue
178
+ for folder in (here, Path(vzd.__file__).parent):
179
+ if (folder / name).is_file():
180
+ return folder / name
181
+ raise ValueError(f"no WAD found for {config}")
182
+
183
+
184
+ class DoomEnvironment:
185
+ """A headless ViZDoom game whose observations are JSON, not pixels."""
186
+
187
+ def __init__(
188
+ self, task: DoomTask, visible: bool = False, record: bool = False, game_args: str = ""
189
+ ):
190
+ """game_args hosts or joins a networked game, such as an arena duel."""
191
+ self.task = task
192
+ self.visible = visible
193
+ self.game_args = game_args
194
+ if task.deathmatch and not game_args:
195
+ # ViZDoom segfaults starting a deathmatch-only map without a network game.
196
+ raise ValueError("deathmatch tasks are played in the arena, not alone")
197
+ # ViZDoom otherwise writes _vizdoom.ini into the working directory.
198
+ self.scratch = tempfile.TemporaryDirectory()
199
+ self.game = None
200
+ self.scenario = None
201
+ self.level = None
202
+ # Every rendered tic while recording, as height x width x RGB arrays.
203
+ self.frames: list | None = [] if record else None
204
+ self.tics = task.tics_per_action
205
+ # A networked game waits for every player at init, and each of its episodes needs a
206
+ # fresh game, so it starts with its first episode instead.
207
+ if not self.game_args:
208
+ self.start(task.scenario)
209
+
210
+ def start(self, scenario: str, map_name: str = "", seed: int | None = None) -> None:
211
+ """Start a game for the scenario; each episode then picks its own map."""
212
+ config = scenario_config(scenario)
213
+ if self.game is not None:
214
+ self.game.close()
215
+ task = self.task
216
+ self.config = config
217
+ self.game = vzd.DoomGame()
218
+ self.game.load_config(config)
219
+ # Graphs take as long as they need per action, so the game always waits for them.
220
+ self.game.set_mode(vzd.Mode.PLAYER)
221
+ if map_name:
222
+ self.game.set_doom_map(map_name)
223
+ if seed is not None:
224
+ self.game.set_seed(seed)
225
+ self.game.add_game_args(self.game_args)
226
+ self.game.set_doom_config_path(os.path.join(self.scratch.name, "_vizdoom.ini"))
227
+ self.game.set_window_visible(self.visible)
228
+ self.game.set_labels_buffer_enabled(True)
229
+ self.game.set_audio_buffer_enabled(False)
230
+ self.game.set_sectors_info_enabled("rangefinder" in task.observe)
231
+ self.game.set_objects_info_enabled(bool({"compass", "target"} & set(task.observe)))
232
+ self.game.set_screen_format(vzd.ScreenFormat.RGB24)
233
+ for name in task.variables:
234
+ self.game.add_available_game_variable(getattr(vzd.GameVariable, name))
235
+ for name, value in task.rewards.items():
236
+ if name not in CUSTOM_REWARDS:
237
+ getattr(self.game, f"set_{name}")(value)
238
+ if task.episode_timeout is not None:
239
+ self.game.set_episode_timeout(task.episode_timeout)
240
+ pressed = {label: set(label.split("+")) for label in task.labels}
241
+ wanted = set().union(*pressed.values())
242
+ unknown = {name for name in wanted if not hasattr(vzd.Button, name)}
243
+ if unknown:
244
+ raise ValueError(f"unknown Doom buttons: {sorted(unknown)}")
245
+ # The labels are the action space, so a scenario's own buttons are only a starting point.
246
+ for name in sorted(wanted.difference(b.name for b in self.game.get_available_buttons())):
247
+ self.game.add_available_button(getattr(vzd.Button, name))
248
+ self.game.init()
249
+ self.scenario = scenario
250
+ buttons = [button.name for button in self.game.get_available_buttons()]
251
+ self.buttons = {label: [b in pressed[label] for b in buttons] for label in task.labels}
252
+ self.variables = [v.name for v in self.game.get_available_game_variables()]
253
+ self.width = self.game.get_screen_width()
254
+ self.height = self.game.get_screen_height()
255
+
256
+ def reset(self, episode: int | Episode) -> dict:
257
+ episode = self.task.episode(episode)
258
+ config = scenario_config(episode.scenario)
259
+ map_name = (episode.map or setting(config, "doom_map") or "").upper()
260
+ if self.game_args:
261
+ # Networked players cannot change maps or restart episodes on their own, so every
262
+ # episode is a new game that everyone joins at init.
263
+ self.start(episode.scenario, map_name, episode.seed)
264
+ else:
265
+ if episode.scenario != self.scenario:
266
+ self.start(episode.scenario)
267
+ if map_name:
268
+ self.game.set_doom_map(map_name)
269
+ self.game.set_seed(episode.seed)
270
+ self.game.new_episode()
271
+ self.level = None
272
+ if "route" in self.task.observe or "progress" in self.task.rewards:
273
+ if not map_name:
274
+ raise ValueError(f"routing needs a map name for {episode.scenario}")
275
+ self.level = load_level(level_file(self.config, map_name), map_name)
276
+ if self.frames is not None:
277
+ self.frames = [self.game.get_state().screen_buffer.copy()]
278
+ self.visited = Counter()
279
+ self.exited = False
280
+ self.recent = deque(maxlen=self.task.memory_steps)
281
+ self.trail = deque(maxlen=self.task.memory_steps + 1)
282
+ self.visit()
283
+ self.plot()
284
+ self.best = self.remaining
285
+ return self.observe()
286
+
287
+ def plot(self) -> None:
288
+ """Recompute the walking distance from the player's position to the exit."""
289
+ self.remaining = None
290
+ if self.level is not None:
291
+ x, y, _ = self.pose()
292
+ self.remaining = self.level.remaining(x, y)
293
+
294
+ def pose(self) -> tuple[float, float, float]:
295
+ variable = self.game.get_game_variable
296
+ GV = vzd.GameVariable
297
+ return variable(GV.POSITION_X), variable(GV.POSITION_Y), variable(GV.ANGLE)
298
+
299
+ def visit(self) -> bool:
300
+ """Count the player's current cell; True when it was never visited before."""
301
+ x, y, _ = self.pose()
302
+ self.trail.append((x, y))
303
+ cell = (math.floor(x / self.task.cell_size), math.floor(y / self.task.cell_size))
304
+ self.visited[cell] += 1
305
+ self.cell = cell
306
+ return self.visited[cell] == 1
307
+
308
+ def step(self, action: str) -> tuple[float, dict | None]:
309
+ if self.frames is None:
310
+ reward = self.game.make_action(self.buttons[action], self.tics)
311
+ else:
312
+ # One tic at a time so every frame is captured; the game advances identically.
313
+ self.game.set_action(self.buttons[action])
314
+ reward = 0.0
315
+ for _ in range(self.tics):
316
+ self.game.advance_action(1)
317
+ reward += self.game.get_last_reward()
318
+ if self.game.is_episode_finished():
319
+ break
320
+ self.frames.append(self.game.get_state().screen_buffer.copy())
321
+ self.recent.append(action)
322
+ # A level exit pays its bonus in the step that ends the episode, and no other event
323
+ # pays anything near it, so the bonus is how an exit is known from a death or timeout.
324
+ if self.game.is_episode_finished():
325
+ self.exited = reward >= self.task.rewards.get("map_exit_reward", math.inf)
326
+ # Only deathmatch outlives a death; single-player episodes end above.
327
+ if not self.game.is_episode_finished() and self.game.is_player_dead():
328
+ self.game.respawn_player()
329
+ # A game with no state has nothing left to decide on: this seat is between lives, or
330
+ # its opponent has left the duel. Either way the episode is over for this player.
331
+ if self.game.is_episode_finished() or self.game.get_state() is None:
332
+ return reward, None
333
+ if self.visit():
334
+ reward += self.task.rewards.get("explore", 0.0)
335
+ self.plot()
336
+ if self.remaining is not None and (self.best is None or self.remaining < self.best):
337
+ if self.best is not None:
338
+ reward += self.task.rewards.get("progress", 0.0) * (self.best - self.remaining)
339
+ self.best = self.remaining
340
+ return reward, self.observe()
341
+
342
+ def observe(self) -> dict:
343
+ state = self.game.get_state()
344
+ objects = [
345
+ {
346
+ "name": label.object_name,
347
+ "category": label.object_category,
348
+ # Screen-relative centers in [-1, 1], left/top negative; rounded for short prompts.
349
+ "x": round((label.x + label.width / 2) / self.width * 2 - 1, 3),
350
+ "y": round((label.y + label.height / 2) / self.height * 2 - 1, 3),
351
+ "width": round(label.width / self.width, 3),
352
+ "height": round(label.height / self.height, 3),
353
+ }
354
+ for label in state.labels
355
+ # "Self" is the player's own weapon sprite.
356
+ if label.object_category != "Self"
357
+ ]
358
+ # ViZDoom reports None, not an empty array, when a scenario has no variables.
359
+ values = state.game_variables if self.variables else []
360
+ variables = dict(zip(self.variables, map(float, values), strict=True))
361
+ observation = {"objects": objects, "variables": variables}
362
+ x, y, angle = self.pose()
363
+ if "target" in self.task.observe:
364
+ target = self.target(state, objects, x, y)
365
+ if target is not None:
366
+ observation["target"] = target
367
+ if "pose" in self.task.observe:
368
+ observation["pose"] = {"x": round(x, 1), "y": round(y, 1), "angle": round(angle, 1)}
369
+ if "rangefinder" in self.task.observe:
370
+ observation["rangefinder"] = self.rangefinder(state, x, y, angle)
371
+ if "compass" in self.task.observe:
372
+ targets = []
373
+ for thing in state.objects:
374
+ if thing.name not in self.task.compass_targets:
375
+ continue
376
+ dx, dy = thing.position_x - x, thing.position_y - y
377
+ # The player's own body stands at its own position and is not a bearing.
378
+ if not dx and not dy:
379
+ continue
380
+ # Degrees from the facing direction; positive is to the left (counterclockwise).
381
+ bearing = (math.degrees(math.atan2(dy, dx)) - angle + 180) % 360 - 180
382
+ targets.append(
383
+ {
384
+ "name": thing.name,
385
+ "bearing": round(bearing, 1),
386
+ "distance": round(math.hypot(dx, dy)),
387
+ }
388
+ )
389
+ targets.sort(key=lambda target: target["distance"])
390
+ # The nearest of each name comes first, so a lone opponent across the map is never
391
+ # crowded out by the pickups underfoot; the rest of the room goes by distance.
392
+ nearest = list({target["name"]: target for target in reversed(targets)}.values())
393
+ nearest.sort(key=lambda target: target["distance"])
394
+ rest = [target for target in targets if target not in nearest]
395
+ observation["compass"] = (nearest + rest)[: self.task.compass_size]
396
+ found = self.level.waypoint(x, y) if "route" in self.task.observe else None
397
+ if found is not None:
398
+ (target_x, target_y), ahead = found
399
+ dx, dy = target_x - x, target_y - y
400
+ observation["route"] = {
401
+ "next": ahead,
402
+ "bearing": round((math.degrees(math.atan2(dy, dx)) - angle + 180) % 360 - 180, 1),
403
+ "distance": round(math.hypot(dx, dy)),
404
+ "remaining": round(self.remaining),
405
+ }
406
+ if "memory" in self.task.observe:
407
+ observation["memory"] = {
408
+ "previous_actions": list(self.recent),
409
+ "distance_moved": round(math.dist(self.trail[0], self.trail[-1]), 1),
410
+ "visits_here": self.visited[self.cell],
411
+ "explored_cells": len(self.visited),
412
+ }
413
+ return observation
414
+
415
+ def target(self, state, objects: list[dict], x: float, y: float) -> dict | None:
416
+ """The nearest live monster or opponent on screen, with its map distance."""
417
+ # Killed monsters and players leave the labels, so every one on screen is alive; the
418
+ # player's own body is category Self and never listed. Labels share
419
+ # IDs with map objects, whose positions give each visible monster's distance.
420
+ places = {thing.id: thing for thing in state.objects}
421
+ visible = [label for label in state.labels if label.object_category != "Self"]
422
+ ranged = [
423
+ (
424
+ math.hypot(
425
+ places[label.object_id].position_x - x, places[label.object_id].position_y - y
426
+ ),
427
+ entry,
428
+ )
429
+ for label, entry in zip(visible, objects, strict=True)
430
+ if entry["category"] in ("Monster", "Player") and label.object_id in places
431
+ ]
432
+ if not ranged:
433
+ return None
434
+ distance, nearest = min(ranged, key=lambda pair: pair[0])
435
+ return {
436
+ "name": nearest["name"],
437
+ "x": nearest["x"],
438
+ "width": nearest["width"],
439
+ "distance": round(distance),
440
+ }
441
+
442
+ def rangefinder(self, state, x: float, y: float, angle: float) -> list[int]:
443
+ """Map-unit distance to the nearest wall along rays across the view, left to right."""
444
+ starts, ends = walls(state)
445
+ span = FIELD_OF_VIEW / 2
446
+ step = FIELD_OF_VIEW / (self.task.rays - 1)
447
+ readings = []
448
+ for index in range(self.task.rays):
449
+ heading = math.radians(angle + span - index * step)
450
+ direction = np.array([math.cos(heading), math.sin(heading)])
451
+ edges = ends - starts
452
+ offset = starts - (x, y)
453
+ denominator = cross(direction, edges)
454
+ with np.errstate(divide="ignore", invalid="ignore"):
455
+ distance = cross(offset, edges) / denominator
456
+ along = cross(offset, direction) / denominator
457
+ hits = distance[(denominator != 0) & (distance > 0) & (along >= 0) & (along <= 1)]
458
+ nearest = min(
459
+ hits.min(initial=self.task.rangefinder_range), self.task.rangefinder_range
460
+ )
461
+ readings.append(round(float(nearest)))
462
+ return readings
463
+
464
+ def stats(self) -> dict:
465
+ """Level progress counters; ViZDoom keeps them readable after the episode ends."""
466
+ variable = self.game.get_game_variable
467
+ GV = vzd.GameVariable
468
+ stats = {
469
+ "kills": int(variable(GV.KILLCOUNT)),
470
+ "items": int(variable(GV.ITEMCOUNT)),
471
+ "secrets": int(variable(GV.SECRETCOUNT)),
472
+ "explored_cells": len(self.visited),
473
+ }
474
+ if self.game_args:
475
+ stats["frags"] = int(variable(GV.FRAGCOUNT))
476
+ stats["deaths"] = int(variable(GV.DEATHCOUNT))
477
+ else:
478
+ stats["exited"] = self.exited
479
+ return stats
480
+
481
+ def close(self) -> None:
482
+ if self.game is not None:
483
+ self.game.close()
484
+ self.scratch.cleanup()
485
+
486
+ def __enter__(self) -> Self:
487
+ return self
488
+
489
+ def __exit__(self, *exc) -> None:
490
+ self.close()
491
+
492
+
493
+ class Episodes:
494
+ """Scores a graph by the normalized reward of one episode per seed."""
495
+
496
+ def __init__(self, task: DoomTask, env: DoomEnvironment):
497
+ self.task = task
498
+ self.env = env
499
+
500
+ def __call__(self, program: DecisionProgram, episode: int | Episode) -> Outcome:
501
+ episode = self.task.episode(episode)
502
+ observation, steps, models, error = self.env.reset(episode), [], set(), None
503
+ while observation is not None:
504
+ prediction = program(state=observation)
505
+ models.update(prediction.models)
506
+ if prediction.undecided:
507
+ error = f"{undecided(program.plan, prediction)}, at step {len(steps)}"
508
+ break
509
+ reward, following = self.env.step(prediction.label)
510
+ steps.append({"observation": observation, "action": prediction.label, "reward": reward})
511
+ observation = following
512
+
513
+ total = sum(step["reward"] for step in steps)
514
+ low, high = self.task.reward_range
515
+ if error is None and not low <= total <= high:
516
+ raise ValueError(f"episode reward {total} is outside reward_range {[low, high]}")
517
+ score = 0.0 if error else (total - low) / (high - low)
518
+ actions = Counter(step["action"] for step in steps)
519
+ stats = self.env.stats()
520
+ record = {
521
+ "seed": episode.seed,
522
+ "map": episode.map,
523
+ "reward": total,
524
+ "score": score,
525
+ "steps": len(steps),
526
+ "actions": dict(actions),
527
+ "stats": stats,
528
+ "error": error,
529
+ }
530
+ feedback = error or (
531
+ f"Episode {episode.seed} on {episode.map or self.task.scenario}: reward {total} in "
532
+ f"range {[low, high]} (score {score:.3f}) after {len(steps)} actions, with {stats}. "
533
+ "Higher reward is better."
534
+ )
535
+ return Outcome(
536
+ score=score,
537
+ record=record,
538
+ trace={
539
+ # The observation itself, not wrapped: search writes read paths by copying them
540
+ # out of this example, so a key around it becomes a prefix on every path it writes.
541
+ "Inputs": steps[0]["observation"] if steps else observation,
542
+ "Generated Outputs": {
543
+ "actions": dict(actions),
544
+ "last_steps": steps[-self.task.feedback_steps :],
545
+ },
546
+ "Feedback": feedback,
547
+ },
548
+ models=sorted(models),
549
+ error=error,
550
+ )
rulesmith/extract.py ADDED
@@ -0,0 +1,77 @@
1
+ """Reading free text into typed fields with a generative model, for rules to test."""
2
+
3
+ import json
4
+ import re
5
+
6
+ import dspy
7
+ from dspy.utils.exceptions import AdapterParseError
8
+
9
+ from rulesmith.graph import Read, State
10
+ from rulesmith.optimize import text_adapter
11
+
12
+ # A field the model wrote as anything but a plain number is not one, and is left absent.
13
+ NUMBER = re.compile(r"-?\d+(\.\d+)?([eE][-+]?\d+)?")
14
+
15
+
16
+ def described(field: Read) -> str:
17
+ if field.categories is not None:
18
+ return f"One of: {', '.join(field.categories)}. None if the record does not say."
19
+ low, high = field.bounds
20
+ return f"A number from {low:g} to {high:g}. None if the record does not say."
21
+
22
+
23
+ class Extractor:
24
+ """A model asked to fill a few fields from a record, each answer kept as text: whether it is
25
+ a number in bounds or one of the categories is the read's to decide, as for any input. Each
26
+ record is read once per set of fields, since search runs a graph on the same examples again
27
+ and again."""
28
+
29
+ def __init__(self, lm: dspy.LM):
30
+ self.lm = lm
31
+ self.found: dict[str, tuple[dict, dict]] = {}
32
+
33
+ def __call__(
34
+ self, instructions: str, fields: dict[str, Read], record: State
35
+ ) -> tuple[dict, dict]:
36
+ key = json.dumps(
37
+ [instructions, {name: described(field) for name, field in fields.items()}, record],
38
+ sort_keys=True,
39
+ )
40
+ if key not in self.found:
41
+ self.found[key] = self.extract(instructions, fields, record)
42
+ return self.found[key]
43
+
44
+ def extract(
45
+ self, instructions: str, fields: dict[str, Read], record: State
46
+ ) -> tuple[dict, dict]:
47
+ signature = dspy.make_signature(
48
+ {"record": (State, dspy.InputField())}
49
+ | {
50
+ name: (str, dspy.OutputField(desc=described(field)))
51
+ for name, field in fields.items()
52
+ },
53
+ instructions,
54
+ )
55
+ usage = {
56
+ "model": self.lm.model,
57
+ "input_tokens": 0,
58
+ "output_tokens": 0,
59
+ "extracted": list(fields),
60
+ }
61
+ with dspy.context(lm=self.lm, adapter=text_adapter(), track_usage=True):
62
+ try:
63
+ predicted = dspy.Predict(signature)(record=record)
64
+ except AdapterParseError as error:
65
+ # A reply with no fields in it fills none: the rules see every field absent.
66
+ return {}, usage | {"error": str(error)}
67
+ for counted in predicted.get_lm_usage().values():
68
+ usage["input_tokens"] += counted.get("prompt_tokens") or 0
69
+ usage["output_tokens"] += counted.get("completion_tokens") or 0
70
+ values = {}
71
+ for name, field in fields.items():
72
+ text = str(getattr(predicted, name)).strip()
73
+ if field.categories is not None:
74
+ values[name] = text
75
+ elif NUMBER.fullmatch(text):
76
+ values[name] = float(text)
77
+ return values, usage