Add observation_function for full custom control over the observation

Lets a game define how its state becomes an observation as a plain
Python function (module:attr), used identically by GameEnvironment
(training) and TrainedPolicy (inference) instead of two independently
maintained encoding paths. Removes the egocentric/egocentric_player/
egocentric_radius flags — cropping is now something an
observation_function does itself by calling egocentric_board(), and
extras_size is discovered from one sampled observation instead of
being configured via observe_state_sizes.
This commit is contained in:
Chris Proctor
2026-06-23 20:45:48 -04:00
parent 426e59a54e
commit 0cd3c3b488
14 changed files with 479 additions and 205 deletions

View File

@@ -7,7 +7,7 @@ Game description
----------------
.. autoclass:: retro_gamer.GameMetadata
:members: from_pyproject, from_dict, validate
:members: from_pyproject, from_dict, validate, resolve_observation_function
Training
--------

View File

@@ -351,6 +351,13 @@ engineering decisions live: what derived quantities should the agent
see, and does giving it those values give it an advantage a human
player would not have?
``character_set``/``observe_state`` cover the common cases, but
sometimes you want full control over how the board becomes numbers — for
example, cropping it to a window centered on the agent rather than always
seeing the whole thing. ``observation_function`` (see :doc:`reference`) lets
you write that transformation as ordinary code instead of a combination of
flags.
Neural network architectures
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~

View File

@@ -100,9 +100,10 @@ matters.
**Observation design** determines what information is available to the
agent. If you leave a character out of the ``character_set``, the agent
will not distinguish it from empty space. If the game module defines a
``get_state()`` function, the agent also receives those computed values
as part of its observation. The consequences of these choices for what
will not distinguish it from empty space. If you list keys in
``observe_state``, the agent also receives those computed values as part
of its observation — or, for full control, an ``observation_function`` can
replace the encoding entirely. The consequences of these choices for what
the agent can learn are reasonably predictable — and making and checking
those predictions is exactly the kind of reasoning the tool is designed
to support.

View File

@@ -135,57 +135,70 @@ or tuples must always have the same length from episode to episode.
Always initialize every observed key with a placeholder of the
correct type and length before the first ``game.step()`` call.
``observe_state_sizes`` (auto-discovered)
.. _observation-function:
``observation_function`` (default: none)
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
A table mapping each ``observe_state`` key to its flat size (``1`` for
scalars, ``N`` for sequences of length N). This is written automatically
to ``config.toml`` the first time ``retro-gamer train`` runs, after the
trainer samples ``game.state`` to discover the actual sizes:
**Optional**, set in ``[metadata]`` (alongside ``actions``/``reward``/
``character_set``/``board_size``), not in ``[preprocessing]``. A
``"module:attr"`` string naming a function ``f(game) -> observation`` that
fully replaces the built-in board/``observe_state`` encoding described
above. Mutually exclusive with ``observe_state`` — they're two conflicting
ways of describing the same thing, and setting both raises an error.
.. code-block:: toml
observe_state_sizes = {board_state = 9}
[metadata]
observation_function = "my_game:get_observation"
You do not need to set this manually. Once written, it is used to
detect changes in state shape when resuming training—an incompatible
change here requires running ``retro-gamer clean`` and starting fresh.
For DQN training, the function must return a flat, numeric, fixed-length
1-D array every time it's called — the same contract the built-in encoder
follows: a flattened one-hot board (sized from ``character_set`` ×
``board_size``, if ``board = true``) followed by any extra features, all in
one vector. ``character_set`` and ``board_size`` stay required either way,
because that's what lets ``observation_function`` also use a spatial
(``spatial = true``) network — the trainer slices the flat vector back into
a board tensor using exactly those two fields, the same way it does for the
built-in encoder. The size of whatever comes after the board (``extras_size``)
is not configured; it's measured automatically from one sampled observation
when training starts.
``egocentric`` (default: ``false``)
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
This is also how you get an egocentric (cropped, player-centered) board now —
there's no longer a built-in flag for it. Call ``egocentric_board()`` and
``encode_board()`` yourself, from :mod:`retro_gamer.observation`, inside your
own function, and declare ``board_size`` to match your crop:
When ``true``, the board observation is cropped to a square window
centred on a specific agent rather than the full board. This gives the
agent a local, first-person-like view and makes the observation
invariant to the agent's absolute position on the board.
.. code-block:: python
Requires ``egocentric_player`` and ``egocentric_radius``.
import numpy as np
from retro.views.headless import HeadlessView
from retro_gamer.observation import egocentric_board, encode_board, encode_state
``egocentric_player``
~~~~~~~~~~~~~~~~~~~~~~
CHARACTER_SET = ["@", "*", ">", "<", "^", "v"]
RADIUS = 8
The name of the agent to use as the centre of the egocentric crop.
Must match the ``name`` attribute of one of the game's agents.
def egocentric_observation(game):
view = HeadlessView()
view.on_game_start(game)
view.render(game)
head = game.get_agent_by_name("Snake head")
cropped = egocentric_board(view.board_characters, head.position, RADIUS)
board_vec = encode_board(cropped, CHARACTER_SET).flatten()
extras = encode_state(game.state, ["apple_dx", "apple_dy"])
return np.concatenate([board_vec, extras])
.. code-block:: toml
egocentric_player = "Snake head"
[metadata]
board_size = [17, 17] # 2*RADIUS + 1
observation_function = "my_module:egocentric_observation"
``egocentric_radius``
~~~~~~~~~~~~~~~~~~~~~~
The half-side-length of the egocentric crop window, in cells. The
resulting observation covers a ``(2r+1) × (2r+1)`` region. Larger
values give the agent a wider view; smaller values focus it on the
immediate vicinity.
.. code-block:: toml
egocentric_radius = 8 # 17×17 window
When ``egocentric_radius`` is set, ``board_size`` in ``[metadata]`` is
automatically updated to ``[2r+1, 2r+1]`` so the network is sized
correctly.
Outside DQN training — for example, BabySnake's tabular Q-learning lab, which
uses :class:`~retro_gamer.GameEnvironment` directly without
:class:`~retro_gamer.DQNTrainer` — there's no 1-D requirement at all.
``observation_function`` can return anything you want to use as your
observation, including a plain tuple used as a dict key.
.. _hyperparameters:
@@ -367,11 +380,11 @@ prints a message and exits immediately. To keep training, increase
unusable. If you change any of the following, ``retro-gamer train`` will
detect the mismatch and refuse to resume, with a clear explanation:
- ``actions``, ``reward``, ``character_set``, ``board_size``
(``[metadata]``) — game description
- ``spatial``, ``board``, ``observe_state``, ``observe_state_sizes``,
``egocentric``, ``egocentric_player``, ``egocentric_radius``
(``[preprocessing]``) — observation encoding
- ``actions``, ``reward``, ``character_set``, ``board_size``,
``observation_function``, ``extras_size`` (``[metadata]``) — game
description and observation shape
- ``spatial``, ``board``, ``observe_state`` (``[preprocessing]``) —
observation encoding
- ``hidden_sizes`` (``[model]``) — network architecture
Run ``retro-gamer clean RUN_DIR`` to remove the old checkpoints and start

View File

@@ -127,9 +127,11 @@ The number of exploration turns is controlled by the
The ``[tool.retro-gamer]`` section describes the game. Preprocessing
options—such as ``spatial`` (whether to use a CNN or MLP, default:
``false``), ``egocentric``, and ``observe_state``—live in the
``[preprocessing]`` section of the generated ``config.toml``. You can
edit them there after running ``retro-gamer create``.
``false``) and ``observe_state``—live in the ``[preprocessing]`` section of
the generated ``config.toml``. You can edit them there after running
``retro-gamer create``. For full control over the observation (for example,
a cropped/egocentric board), write an ``observation_function`` instead — see :ref:`observation-function` in the
reference docs for details.
``observe_state``
~~~~~~~~~~~~~~~~~
@@ -408,15 +410,14 @@ checkpoints remain valid:
game or the shape of the network. The saved model weights are
incompatible with the new configuration:
- ``actions``, ``reward``, ``character_set``, ``board_size``
(``[metadata]``) — These define what the agent perceives and what it
can do. Changing them changes the size of the network's input or
output layers; the existing weights no longer fit.
- ``spatial``, ``board``, ``observe_state``, ``observe_state_sizes``,
``egocentric``, ``egocentric_player``, ``egocentric_radius``
(``[preprocessing]``) — These control how the observation is
constructed. Any change here alters the input shape or meaning and
makes existing weights invalid.
- ``actions``, ``reward``, ``character_set``, ``board_size``,
``observation_function``, ``extras_size`` (``[metadata]``) — These define
what the agent perceives and what it can do. Changing them changes the
size of the network's input or output layers; the existing weights no
longer fit.
- ``spatial``, ``board``, ``observe_state`` (``[preprocessing]``) — These
control how the observation is constructed. Any change here alters the
input shape or meaning and makes existing weights invalid.
- ``hidden_sizes`` (``[model]``) — This defines the network's hidden
layers. Changing it changes the shape of the network; the existing
weights no longer fit.

View File

@@ -1,6 +1,6 @@
[project]
name = "retro-gamer"
version = "0.1.1"
version = "0.2.0"
description = "A toolkit for learning reinforcement learning by training agents to play retro games"
readme = "README.md"
requires-python = ">=3.11"

View File

@@ -13,6 +13,12 @@ from retro_gamer.trainer import DQNTrainer, DEFAULTS, MODEL_KEYS
@click.group()
def cli():
"""Train and run RL agents for retro games."""
# Running the installed console script puts its own directory on sys.path,
# not the caller's cwd — but observation_function/game modules are
# typically plain files in the directory you ran retro-gamer from.
cwd = str(Path.cwd())
if cwd not in sys.path:
sys.path.insert(0, cwd)
# ---------------------------------------------------------------------------
@@ -79,8 +85,9 @@ def create(game, output, **hyperparams):
raise click.ClickException(str(e))
game_factory = _load_factory(game_config)
g = game_factory()
metadata.board_size = g.board_size
if metadata.board_size is None:
g = game_factory()
metadata.board_size = g.board_size
metadata.validate()
@@ -108,6 +115,8 @@ def create(game, output, **hyperparams):
else:
click.echo(f" characters : (will be auto-discovered during training)")
click.echo(f" architecture: {'CNN (spatial)' if metadata.spatial else 'MLP (non-spatial)'}")
if metadata.observation_function:
click.echo(f" observation : {metadata.observation_function} (custom)")
# ---------------------------------------------------------------------------

View File

@@ -1,6 +1,5 @@
from __future__ import annotations
import random
import numpy as np
from typing import Callable
from retro.input import ProgrammaticInput
from retro.views.headless import HeadlessView
@@ -9,34 +8,49 @@ from retro_gamer.observation import encode_observation
class GameEnvironment:
"""Gym-style wrapper around a retro game for RL training."""
"""Gym-style wrapper around a retro game for RL training.
The observation returned by reset()/step() comes from one of two mutually
exclusive paths: metadata.observation_function, if set, fully replaces the
built-in board/observe_state encoding — see GameMetadata for its contract.
Otherwise the built-in encoder (encode_observation) is used, configured by
observe_state/board/observe_state_sizes below.
"""
def __init__(
self,
game_factory: Callable,
metadata: GameMetadata,
observe_state: list[str] | None = None,
egocentric: bool = False,
egocentric_player: str | None = None,
egocentric_radius: int | None = None,
board: bool = True,
observe_state_sizes: dict[str, int] | None = None,
):
self.game_factory = game_factory
self.metadata = metadata
self.observe_state = observe_state or []
self.egocentric = egocentric
self.egocentric_player = egocentric_player
self.egocentric_radius = egocentric_radius
self.board = board
self.observe_state_sizes = observe_state_sizes or {}
self._observation_fn = metadata.resolve_observation_function()
if self._observation_fn is not None and self.observe_state:
raise ValueError(
"Both metadata.observation_function and [preprocessing].observe_state "
"are set, but they're two conflicting ways of describing the\n"
"observation. Use observation_function for full custom control, or\n"
"observe_state (with the built-in board encoder) but not both."
)
self.game = None
self.view: HeadlessView | None = None
self.inp: ProgrammaticInput | None = None
self._prev_reward: float = 0.0
def reset(self) -> np.ndarray:
"""Create a fresh game episode and return the initial observation."""
def reset(self):
"""Create a fresh game episode and return the initial observation.
The observation's type depends on metadata.observation_function: a
numpy array when using the built-in encoder (or a custom function
built for DQN training), but it can be anything a custom function
returns — e.g. a plain tuple for tabular use.
"""
self.inp = ProgrammaticInput()
self.view = HeadlessView()
self.game = self.game_factory()
@@ -46,7 +60,7 @@ class GameEnvironment:
self._prev_reward = float(self.game.state.get(self.metadata.reward, 0))
return self._observe()
def step(self, action: str | None) -> tuple[np.ndarray, float, bool]:
def step(self, action: str | None) -> tuple:
"""Advance one turn. Returns (observation, reward, done)."""
self.inp.press(action)
self.game.step()
@@ -55,48 +69,18 @@ class GameEnvironment:
done = not self.game.playing
return obs, reward, done
def _observe(self) -> np.ndarray:
def _observe(self):
if self._observation_fn is not None:
return self._observation_fn(self.game)
state = dict(self.game.state)
if self.observe_state_sizes:
self._check_state_sizes(state)
player_pos = None
if self.egocentric and self.egocentric_player:
agent = self.game.get_agent_by_name(self.egocentric_player)
if agent is not None:
player_pos = agent.position
return encode_observation(
self.view.board_characters,
state,
self.metadata,
self.observe_state,
player_pos=player_pos,
egocentric_radius=self.egocentric_radius,
board=self.board,
)
def _check_state_sizes(self, state: dict):
for key, expected in self.observe_state_sizes.items():
val = state.get(key)
if val is None:
actual = 0
elif isinstance(val, (list, tuple)):
actual = len(val)
else:
actual = 1
if actual != expected:
raise ValueError(
f"State key '{key}' changed size during training:\n"
f" Expected : {expected} (discovered at training start)\n"
f" Got : {actual}\n\n"
f"This means game.state['{key}'] has a different length in some\n"
f"episodes than it had when training started. The neural network\n"
f"has a fixed input size and cannot adapt to changing state shapes.\n\n"
f"Fix: make sure create_game() always initializes '{key}' with a\n"
f"fixed-length value before the game starts each episode.\n"
f"For example, if '{key}' is a list of 9 values, it must always be\n"
f"a list of exactly 9 values — never more, never fewer, never missing."
)
def _delta_reward(self) -> float:
current = float(self.game.state.get(self.metadata.reward, 0))
delta = current - self._prev_reward

View File

@@ -4,6 +4,7 @@ import tomllib
import tomli_w
from dataclasses import dataclass, field
from pathlib import Path
from typing import Callable
@dataclass
@@ -11,9 +12,15 @@ class GameMetadata:
"""Describes a retro game for training purposes.
Required fields: actions, reward.
Optional fields: character_set, spatial.
Discovered fields: board_size (from game.board_size), extras_size (from
the observe_state list in [preprocessing]).
Optional fields: character_set, spatial, observation_function.
Discovered fields: board_size (from game.board_size), extras_size
(computed by DQNTrainer from one sampled observation — never set in a
game's own pyproject.toml).
observation_function, if set, is a "module:attr" string naming a function
``f(game) -> Any`` that fully replaces the built-in board/observe_state
encoding. It is mutually exclusive with the [preprocessing] observe_state
option. See GameEnvironment for how the two paths are selected.
"""
actions: list[str]
reward: str
@@ -21,6 +28,7 @@ class GameMetadata:
spatial: bool = False
board: bool = True
board_size: tuple[int, int] | None = None
observation_function: str | None = None
extras_size: int = 0
def validate(self):
@@ -59,6 +67,48 @@ class GameMetadata:
"If you're not sure what characters your game uses, remove character_set\n"
"entirely and the trainer will discover them automatically."
)
if self.observation_function is not None:
if not isinstance(self.observation_function, str) or ':' not in self.observation_function:
raise ValueError(
f"'observation_function' must be a string of the form 'module:attr', "
f"but got: {self.observation_function!r}\n"
"Example: observation_function = \"my_game:get_observation\"\n"
"This should name a function f(game) -> observation that fully\n"
"describes what your agent observes each turn."
)
def resolve_observation_function(self) -> Callable | None:
"""Import and return the function named by observation_function, or None if unset.
Raises ValueError with an actionable message if the string isn't
"module:attr", the module can't be imported, or it has no such attribute.
"""
if self.observation_function is None:
return None
if ':' not in self.observation_function:
raise ValueError(
f"'observation_function' must be of the form 'module:attr', but got "
f"{self.observation_function!r} (no ':' found).\n"
"Example: observation_function = \"my_game:get_observation\""
)
module_name, attr_name = self.observation_function.split(':', 1)
try:
module = importlib.import_module(module_name)
except ImportError as e:
raise ValueError(
f"Could not import module {module_name!r} for observation_function "
f"{self.observation_function!r}: {e}\n"
"Make sure the module is importable (e.g. on PYTHONPATH or installed)."
) from e
try:
return getattr(module, attr_name)
except AttributeError:
raise ValueError(
f"Module {module_name!r} has no attribute {attr_name!r} "
f"(from observation_function = {self.observation_function!r}).\n"
f"Define a function named '{attr_name}' in {module_name} that takes a "
"game instance and returns its observation."
) from None
@classmethod
def from_pyproject(cls, module_name: str) -> GameMetadata:
@@ -102,17 +152,22 @@ class GameMetadata:
character_set=d.get('character_set'),
spatial=d.get('spatial', False),
board_size=board_size,
observation_function=d.get('observation_function'),
extras_size=d.get('extras_size', 0),
)
def to_dict(self) -> dict:
d = {
'actions': self.actions,
'reward': self.reward,
'extras_size': self.extras_size,
}
if self.board_size is not None:
d['board_size'] = list(self.board_size)
if self.character_set is not None:
d['character_set'] = self.character_set
if self.observation_function is not None:
d['observation_function'] = self.observation_function
return d
def to_toml(self, path: str | Path):

View File

@@ -45,17 +45,9 @@ class TrainedPolicy:
pre = config.get('preprocessing', {})
self._metadata.spatial = pre.get('spatial', False)
self._metadata.board = pre.get('board', True)
observe_state_sizes = pre.get('observe_state_sizes', {})
self._observe_state: list[str] = pre.get('observe_state', [])
self._egocentric: bool = pre.get('egocentric', False)
self._egocentric_player: str | None = pre.get('egocentric_player')
self._egocentric_radius: int | None = pre.get('egocentric_radius')
self._board: bool = pre.get('board', True)
if observe_state_sizes:
self._metadata.extras_size = sum(observe_state_sizes.values())
else:
self._metadata.extras_size = len(self._observe_state)
self._observation_fn = self._metadata.resolve_observation_function()
hyperparams = {**DEFAULTS, **config.get('model', {}), **config.get('training', {})}
self._model, _ = build_network(self._metadata, hyperparams)
@@ -78,26 +70,19 @@ class TrainedPolicy:
def get_action(self, game) -> str | None:
"""Return the key the model recommends this turn, or None for no-op."""
view = HeadlessView()
view.on_game_start(game)
view.render(game)
board_chars = view.board_characters
player_pos = None
if self._egocentric and self._egocentric_player:
agent = game.get_agent_by_name(self._egocentric_player)
if agent is not None:
player_pos = agent.position
obs = encode_observation(
board_chars,
dict(game.state),
self._metadata,
self._observe_state,
player_pos=player_pos,
egocentric_radius=self._egocentric_radius,
board=self._board,
)
if self._observation_fn is not None:
obs = self._observation_fn(game)
else:
view = HeadlessView()
view.on_game_start(game)
view.render(game)
obs = encode_observation(
view.board_characters,
dict(game.state),
self._metadata,
self._observe_state,
board=self._board,
)
device = next(self._model.parameters()).device
state_t = torch.as_tensor(obs, dtype=torch.float32).unsqueeze(0).to(device)

View File

@@ -64,25 +64,24 @@ def encode_observation(
state: dict,
metadata: GameMetadata,
observe_state: list[str],
player_pos: tuple[int, int] | None = None,
egocentric_radius: int | None = None,
board: bool = True,
) -> np.ndarray:
"""Encode board and/or selected state values into a flat 1D observation vector.
When *board* is True the board is encoded and prepended to the vector. If
player_pos and egocentric_radius are given the board is first cropped to a
(2r+1)×(2r+1) window centred on the player. For spatial games the board is
encoded channel-first (C, H, W) then flattened; for non-spatial games it is
encoded (H, W, C) then flattened. The state vector is appended at the end.
When *board* is True the board is encoded and prepended to the vector. For
spatial games the board is encoded channel-first (C, H, W) then flattened;
for non-spatial games it is encoded (H, W, C) then flattened. The state
vector is appended at the end.
When *board* is False only the observe_state features are returned.
For cropped (egocentric) boards, write a custom observation_function that
calls egocentric_board() and encode_board() directly instead of using this
function.
"""
if board:
if not metadata.character_set:
raise ValueError("character_set must be set before encoding observations")
if player_pos is not None and egocentric_radius is not None:
board_chars = egocentric_board(board_chars, player_pos, egocentric_radius)
board_enc = encode_board(board_chars, metadata.character_set) # (H, W, C)
if metadata.spatial:
board_vec = board_enc.transpose(2, 0, 1).flatten()

View File

@@ -50,19 +50,17 @@ def _get_device() -> torch.device:
# Fields that make an existing checkpoint incompatible with the current config.
# Changing any of these requires starting training from scratch.
_INCOMPATIBLE_METADATA = {
'actions': 'the list of actions the agent can take (changes output layer size)',
'reward': 'the reward signal — Q-values trained on the old signal are meaningless for the new one',
'character_set': 'the set of board characters (changes input layer size)',
'board_size': 'the board dimensions (changes input layer size)',
'actions': 'the list of actions the agent can take (changes output layer size)',
'reward': 'the reward signal — Q-values trained on the old signal are meaningless for the new one',
'character_set': 'the set of board characters (changes input layer size)',
'board_size': 'the board dimensions (changes input layer size)',
'observation_function': 'how the observation is computed (changes input representation)',
'extras_size': 'the size of the non-board portion of the observation (changes input layer size)',
}
_INCOMPATIBLE_PREPROCESSING = {
'spatial': 'spatial vs non-spatial network type (changes network architecture)',
'board': 'whether the board is included in the observation (changes input size)',
'observe_state': 'the state keys included in the observation (changes input size)',
'observe_state_sizes': 'the size of each observed state key (changes input layer size)',
'egocentric': 'egocentric board transformation (changes input representation)',
'egocentric_player': 'the agent used as the egocentric center (changes input representation)',
'egocentric_radius': 'the egocentric crop radius (changes input layer size)',
}
_INCOMPATIBLE_ARCH = {
'hidden_sizes': 'the hidden layer sizes (changes network shape)',
@@ -274,11 +272,7 @@ class DQNTrainer:
pre = preprocessing or {}
self.observe_state: list[str] = pre.get('observe_state', [])
self.egocentric: bool = pre.get('egocentric', False)
self.egocentric_player: str | None = pre.get('egocentric_player', None)
self.egocentric_radius: int | None = pre.get('egocentric_radius', None)
self.board: bool = pre.get('board', True)
self.observe_state_sizes: dict[str, int] = pre.get('observe_state_sizes', {})
if self.board is False and metadata.spatial:
raise ValueError(
@@ -286,18 +280,11 @@ class DQNTrainer:
"A CNN requires a 2-D board to operate on. Either set spatial = false\n"
"or keep board = true."
)
if self.board is False and not self.observe_state:
if self.board is False and not self.observe_state and metadata.observation_function is None:
raise ValueError(
"preprocessing.board = false requires at least one entry in observe_state.\n"
"With board=false, the agent observes only the game state variables listed\n"
"in observe_state — if that list is empty, there is nothing to observe."
)
if self.egocentric and not self.egocentric_radius:
raise ValueError(
"preprocessing.egocentric = true requires egocentric_radius.\n"
"Choose a value based on how far the agent needs to see, e.g.:\n"
" egocentric_radius = 5 # 11×11 tight local view\n"
" egocentric_radius = 8 # 17×17 wider view"
"preprocessing.board = false requires at least one entry in observe_state\n"
"(or a metadata.observation_function). With board=false and no\n"
"observe_state, there is nothing for the agent to observe."
)
metadata.board = self.board
@@ -306,28 +293,16 @@ class DQNTrainer:
g = game_factory()
metadata.board_size = g.board_size
if self.egocentric_radius:
side = 2 * self.egocentric_radius + 1
metadata.board_size = (side, side)
self.env = GameEnvironment(
game_factory, metadata,
observe_state=self.observe_state,
egocentric=self.egocentric,
egocentric_player=self.egocentric_player,
egocentric_radius=self.egocentric_radius,
board=self.board,
observe_state_sizes=self.observe_state_sizes,
)
if metadata.character_set is None and self.board:
self._discover_character_set()
if self.observe_state and not self.observe_state_sizes:
self._discover_observe_state_sizes()
self.env.observe_state_sizes = self.observe_state_sizes
metadata.extras_size = sum(self.observe_state_sizes.values()) if self.observe_state_sizes else 0
self._discover_extras_size()
self.device = _get_device()
@@ -467,6 +442,7 @@ class DQNTrainer:
def _run_episode(self) -> tuple[float, int, float, bool]:
state = self.env.reset()
self._check_obs_length(state)
total_reward = 0.0
total_loss = 0.0
loss_count = 0
@@ -477,6 +453,7 @@ class DQNTrainer:
action_key = self._idx_to_key(action_idx)
next_state, reward, done = self.env.step(action_key)
self._check_obs_length(next_state)
self.memory.push(state, action_idx, reward, next_state, done)
if self.total_steps % self.hp['train_every'] == 0:
@@ -565,13 +542,9 @@ class DQNTrainer:
return {
'metadata': self.metadata.to_dict(),
'preprocessing': {
'spatial': self.metadata.spatial,
'board': self.board,
'observe_state': self.observe_state,
'observe_state_sizes': self.observe_state_sizes,
'egocentric': self.egocentric,
'egocentric_player': self.egocentric_player,
'egocentric_radius': self.egocentric_radius,
'spatial': self.metadata.spatial,
'board': self.board,
'observe_state': self.observe_state,
},
'hidden_sizes': self.hp['hidden_sizes'],
}
@@ -636,15 +609,61 @@ class DQNTrainer:
f"after {self.hp['exploration_turns']} exploration turns: {chars}"
)
def _discover_observe_state_sizes(self):
"""Sample game.state to determine the flat size of each observe_state key."""
self.env.reset()
state = dict(self.env.game.state)
sizes = {}
for key in self.observe_state:
val = state.get(key, 0)
sizes[key] = len(val) if isinstance(val, (list, tuple)) else 1
self.observe_state_sizes = sizes
def _discover_extras_size(self):
"""Sample one observation to determine extras_size (everything past the board).
This is also where the observation contract is enforced for training:
GameEnvironment itself has no opinion about what observations look like
(BabySnake, for example, uses it with a plain tuple), but DQNTrainer
needs a flat, numeric, fixed-length vector to feed a neural network.
"""
sample = self.env.reset()
try:
arr = np.asarray(sample, dtype=np.float32)
except (TypeError, ValueError) as e:
raise ValueError(
"Could not convert the observation to a numeric array for training:\n"
f" {sample!r}\n\n"
"DQNTrainer requires observation_function (or the built-in encoder) to\n"
"return a flat, numeric array-like value — a list/tuple of numbers or a\n"
"numpy array."
) from e
if arr.ndim != 1:
raise ValueError(
f"Expected a 1-D observation, but got shape {arr.shape}.\n"
"DQNTrainer always works with a single flat vector — board and extras\n"
"(if any) must be combined into one vector before being returned."
)
board_length = 0
if self.board:
C = len(self.metadata.character_set) if self.metadata.character_set else 0
bw, bh = self.metadata.board_size
board_length = C * bw * bh
if len(arr) < board_length:
raise ValueError(
f"The observation has length {len(arr)}, but character_set "
f"({C} chars) x board_size ({bw}x{bh}) = {board_length} is larger "
"than that.\n"
"Check that character_set/board_size match what your observation\n"
"actually encodes, or set board = false if there's no board in it."
)
self.metadata.extras_size = len(arr) - board_length
self._obs_len = len(arr)
def _check_obs_length(self, obs):
"""Raise a friendly error if an observation's length differs from the one discovered at init."""
length = len(np.asarray(obs, dtype=np.float32))
if length != self._obs_len:
raise ValueError(
"Observation length changed during training:\n"
f" Expected : {self._obs_len} (discovered at training start)\n"
f" Got : {length}\n\n"
"The neural network has a fixed input size and cannot adapt to a\n"
"changing observation shape. Make sure observation_function (or the\n"
"game's state) always produces the same length every episode."
)
def _save_config(self):
config_path = self.run_dir / 'config.toml'
@@ -658,13 +677,6 @@ class DQNTrainer:
pre['spatial'] = self.metadata.spatial
pre['board'] = self.board
pre['observe_state'] = self.observe_state
if self.observe_state_sizes:
pre['observe_state_sizes'] = self.observe_state_sizes
pre['egocentric'] = self.egocentric
if self.egocentric_player:
pre['egocentric_player'] = self.egocentric_player
if self.egocentric_radius:
pre['egocentric_radius'] = self.egocentric_radius
config['model'] = {k: v for k, v in self.hp.items() if k in MODEL_KEYS}
config['training'] = {k: v for k, v in self.hp.items() if k not in MODEL_KEYS}
with open(config_path, 'wb') as f:

View File

@@ -0,0 +1,208 @@
"""Tests for the observation_function machinery: GameMetadata resolution,
GameEnvironment delegation, and DQNTrainer's extras_size discovery.
Run with: uv run python -m unittest tests/test_observation_function.py -v
"""
import tempfile
from pathlib import Path
from unittest import TestCase, main
import numpy as np
from retro.game import Game
from retro_gamer.metadata import GameMetadata
from retro_gamer.observation import encode_observation
from retro_gamer.env import GameEnvironment
from retro_gamer.trainer import DQNTrainer
ACTIONS = ["KEY_RIGHT", "KEY_UP", "KEY_LEFT", "KEY_DOWN"]
def sample_observation_fn(game):
"""A trivial observation_function used to test dotted-path resolution."""
return "observation-from-sample_observation_fn"
def tuple_observation_fn(game):
"""Returns a plain tuple — exercises the non-array tabular use case."""
return (1, 2, 3)
def make_game_factory(board_size=(2, 2)):
def factory():
return Game([], {'reward': 0.0, 'score': 0}, board_size=board_size, show_state=False)
return factory
def obs_with_extras(game):
"""Board (character_set=2, board_size=(2,2) -> 8) + 3 extras = 11."""
return np.zeros(2 * 2 * 2 + 3, dtype=np.float32)
def obs_2d(game):
"""Not 1-D — used to test the training-time shape validation."""
return np.zeros((2, 2), dtype=np.float32)
def obs_too_short(game):
"""Shorter than character_set (2) x board_size (2x2) = 8."""
return np.zeros(5, dtype=np.float32)
_changing_obs_counter = {"n": 0}
def changing_obs(game):
"""Returns length 8 on the first call, length 9 on every call after."""
_changing_obs_counter["n"] += 1
length = 8 if _changing_obs_counter["n"] == 1 else 9
return np.zeros(length, dtype=np.float32)
class TestResolveObservationFunction(TestCase):
def test_returns_none_when_unset(self):
metadata = GameMetadata(actions=ACTIONS, reward="reward")
self.assertIsNone(metadata.resolve_observation_function())
def test_resolves_valid_dotted_path(self):
metadata = GameMetadata(
actions=ACTIONS, reward="reward",
observation_function=f"{__name__}:sample_observation_fn",
)
fn = metadata.resolve_observation_function()
self.assertIs(fn, sample_observation_fn)
def test_raises_on_malformed_string(self):
metadata = GameMetadata(actions=ACTIONS, reward="reward", observation_function="no_colon_here")
with self.assertRaises(ValueError):
metadata.resolve_observation_function()
def test_raises_on_unimportable_module(self):
metadata = GameMetadata(
actions=ACTIONS, reward="reward",
observation_function="this_module_does_not_exist:fn",
)
with self.assertRaises(ValueError):
metadata.resolve_observation_function()
def test_raises_on_missing_attribute(self):
metadata = GameMetadata(
actions=ACTIONS, reward="reward",
observation_function=f"{__name__}:no_such_function",
)
with self.assertRaises(ValueError):
metadata.resolve_observation_function()
def test_validate_raises_on_malformed_observation_function(self):
metadata = GameMetadata(actions=ACTIONS, reward="reward", observation_function="no_colon_here")
with self.assertRaises(ValueError):
metadata.validate()
def test_validate_passes_with_well_formed_observation_function(self):
metadata = GameMetadata(
actions=ACTIONS, reward="reward",
observation_function=f"{__name__}:sample_observation_fn",
)
metadata.validate() # should not raise
class TestEncodeObservation(TestCase):
def setUp(self):
self.metadata = GameMetadata(
actions=ACTIONS, reward="reward",
character_set=["@", "*"], board_size=(2, 2),
)
self.board_chars = [["@", " "], [" ", "*"]]
def test_board_plus_extras(self):
obs = encode_observation(
self.board_chars, {"x": 1.0, "y": 2.0}, self.metadata, ["x", "y"],
)
self.assertEqual(obs.shape, (2 * 2 * 2 + 2,))
def test_board_false_returns_only_extras(self):
obs = encode_observation(
self.board_chars, {"x": 1.0, "y": 2.0}, self.metadata, ["x", "y"], board=False,
)
np.testing.assert_array_equal(obs, [1.0, 2.0])
class TestGameEnvironment(TestCase):
def test_custom_observation_function_returns_its_value_directly(self):
metadata = GameMetadata(
actions=ACTIONS, reward="reward",
observation_function=f"{__name__}:tuple_observation_fn",
)
env = GameEnvironment(make_game_factory(), metadata)
self.assertEqual(env.reset(), (1, 2, 3))
obs, reward, done = env.step(None)
self.assertEqual(obs, (1, 2, 3))
def test_mutual_exclusivity_raises(self):
metadata = GameMetadata(
actions=ACTIONS, reward="reward",
observation_function=f"{__name__}:tuple_observation_fn",
)
with self.assertRaises(ValueError):
GameEnvironment(make_game_factory(), metadata, observe_state=["score"])
def test_built_in_path_still_works(self):
metadata = GameMetadata(
actions=ACTIONS, reward="reward",
character_set=["@", "*"], board_size=(2, 2),
)
env = GameEnvironment(make_game_factory(), metadata, observe_state=["score"])
obs = env.reset()
self.assertEqual(obs.shape, (2 * 2 * 2 + 1,))
class TestDQNTrainerExtrasSize(TestCase):
def _trainer(self, metadata, preprocessing=None, board_size=(2, 2)):
tmp = tempfile.TemporaryDirectory()
self.addCleanup(tmp.cleanup)
return DQNTrainer(
make_game_factory(board_size), metadata, Path(tmp.name),
preprocessing=preprocessing,
)
def test_computes_extras_size_from_custom_function(self):
metadata = GameMetadata(
actions=ACTIONS, reward="reward",
character_set=["@", "*"], board_size=(2, 2),
observation_function=f"{__name__}:obs_with_extras",
)
trainer = self._trainer(metadata)
self.assertEqual(trainer.metadata.extras_size, 3)
def test_raises_on_non_1d_observation(self):
metadata = GameMetadata(
actions=ACTIONS, reward="reward",
character_set=["@", "*"], board_size=(2, 2),
observation_function=f"{__name__}:obs_2d",
)
with self.assertRaises(ValueError):
self._trainer(metadata)
def test_raises_when_observation_too_short_for_declared_board(self):
metadata = GameMetadata(
actions=ACTIONS, reward="reward",
character_set=["@", "*"], board_size=(2, 2),
observation_function=f"{__name__}:obs_too_short",
)
with self.assertRaises(ValueError):
self._trainer(metadata)
def test_raises_when_observation_length_changes_mid_run(self):
_changing_obs_counter["n"] = 0
metadata = GameMetadata(
actions=ACTIONS, reward="reward",
observation_function=f"{__name__}:changing_obs",
)
trainer = self._trainer(metadata, preprocessing={'board': False})
self.assertEqual(trainer.metadata.extras_size, 8)
with self.assertRaises(ValueError):
trainer._run_episode()
if __name__ == '__main__':
main()

2
uv.lock generated
View File

@@ -1154,7 +1154,7 @@ wheels = [
[[package]]
name = "retro-gamer"
version = "0.1.1"
version = "0.2.0"
source = { editable = "." }
dependencies = [
{ name = "click" },