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