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:
@@ -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
|
||||
--------
|
||||
|
||||
@@ -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
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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)")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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:
|
||||
|
||||
208
tests/test_observation_function.py
Normal file
208
tests/test_observation_function.py
Normal 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()
|
||||
Reference in New Issue
Block a user