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:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user