From c89609fe771b0b58efe6eeba2e8154f8363a0000 Mon Sep 17 00:00:00 2001 From: Chris Proctor Date: Fri, 26 Jun 2026 06:55:30 -0400 Subject: [PATCH] Refactor CLI/interface: init command, factory field, remove extras_size MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Rename `retro-gamer create` to `retro-gamer init` with positional args (GAME OUTPUT) instead of --game/--output flags - Add [tool.retro-gamer].factory = "module:attr" to declare the game factory function from pyproject.toml instead of relying on a hard-coded create_game attribute - Remove extras_size from user-facing config; it is now measured automatically from a sample observation and never declared - Update docs throughout: create→init, runs/→training/, add factory field documentation, clarify [model] vs [training] hyperparameter sections, remove stale version-history explanations - Bump version to 0.3.0; require retro-games>=2.5.0 --- Makefile | 2 +- docs/api.rst | 2 +- docs/index.rst | 2 +- docs/integration.rst | 21 ++++--- docs/introduction.rst | 9 +-- docs/reference.rst | 119 +++++++++++++++++++++++----------------- docs/walkthrough.rst | 95 ++++++++++++++++---------------- pyproject.toml | 4 +- retro_gamer/cli.py | 44 +++++++++++---- retro_gamer/metadata.py | 48 ++++++++++++++-- retro_gamer/trainer.py | 1 - uv.lock | 6 +- 12 files changed, 217 insertions(+), 136 deletions(-) diff --git a/Makefile b/Makefile index 6f1e91b..e4b414e 100644 --- a/Makefile +++ b/Makefile @@ -4,7 +4,7 @@ CF_DISTRIBUTION = EPA6NHZ2LEH1A .PHONY: build deploy clean build: - uv run --group documentation $(MAKE) -C docs html + uv run --group documentation sphinx-build -M html docs docs/_build deploy: build aws s3 sync docs/_build/html $(S3_BUCKET) diff --git a/docs/api.rst b/docs/api.rst index 3abac9d..6ee2946 100644 --- a/docs/api.rst +++ b/docs/api.rst @@ -7,7 +7,7 @@ Game description ---------------- .. autoclass:: retro_gamer.GameMetadata - :members: from_pyproject, from_dict, validate, resolve_observation_function + :members: from_pyproject, from_dict, validate, resolve_factory, resolve_observation_function Training -------- diff --git a/docs/index.rst b/docs/index.rst index c6bdfae..0188f61 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -58,7 +58,7 @@ Verify the installation by checking the command-line tool: Train and run RL agents for retro games. Commands: - create Create a new training run directory with config.toml. + init Initialize a new training run directory with config.toml. info Print a summary of a training run. play Watch a trained agent play the game. train Train (or resume training) a DQN agent. diff --git a/docs/integration.rst b/docs/integration.rst index 936e790..cf21755 100644 --- a/docs/integration.rst +++ b/docs/integration.rst @@ -19,14 +19,14 @@ Both approaches start by creating a :class:`retro_gamer.TrainedPolicy`: from retro_gamer import TrainedPolicy - ai = TrainedPolicy("runs/snake/") + ai = TrainedPolicy("training/snake/") This reads ``config.toml``, rebuilds the network, and loads the latest checkpoint. To load a specific checkpoint instead: .. code-block:: python - ai = TrainedPolicy("runs/snake/", checkpoint="ep_0500") + ai = TrainedPolicy("training/snake/", checkpoint="ep_0500") PolicyInput: model as player ---------------------------- @@ -40,7 +40,7 @@ it to ``game.play()`` and everything else works exactly as usual: from retro.examples.snake import create_game from retro_gamer import TrainedPolicy, PolicyInput - ai = TrainedPolicy("runs/snake/") + ai = TrainedPolicy("training/snake/") game = create_game() game.play(input_source=PolicyInput(ai, game)) @@ -62,7 +62,7 @@ loaded from disk once — not once per episode. from retro.examples.snake import Apple, SnakeHead from retro_gamer import TrainedPolicy - _ai = TrainedPolicy("runs/snake/") + _ai = TrainedPolicy("training/snake/") class AISnake(SnakeHead): def handle_keystroke(self, k, game): pass # ignore keyboard @@ -112,9 +112,8 @@ from the game state. To train an enemy: .. code-block:: console - % retro-gamer create --game my_game:create_enemy_training_game \ - --output runs/enemy/ - % retro-gamer train runs/enemy/ + % retro-gamer init games/my_game training/enemy + % retro-gamer train training/enemy/ 3. **Embed the trained model in your main game** using ``get_action``, exactly as shown above. @@ -145,7 +144,7 @@ once per episode: # enemy_training_game.py from retro_gamer import TrainedPolicy - _player = TrainedPolicy("runs/player/") # loaded once when the module is imported + _player = TrainedPolicy("training/player/") # loaded once when the module is imported def create_game(): enemy = EnemyAgent() @@ -156,9 +155,9 @@ You then alternate training runs: .. code-block:: console - % retro-gamer train runs/player/ # train player against current enemy - % retro-gamer train runs/enemy/ # train enemy against updated player - % retro-gamer train runs/player/ # train player again + % retro-gamer train training/player/ # train player against current enemy + % retro-gamer train training/enemy/ # train enemy against updated player + % retro-gamer train training/player/ # train player again # ... How many episodes to run before switching is itself a design decision: too diff --git a/docs/introduction.rst b/docs/introduction.rst index 012ac1d..1b0e74d 100644 --- a/docs/introduction.rst +++ b/docs/introduction.rst @@ -58,6 +58,7 @@ A typical workflow looks like this. First, describe your game in the .. code-block:: toml [tool.retro-gamer] + factory = "snake:create_game" actions = ["KEY_RIGHT", "KEY_UP", "KEY_LEFT", "KEY_DOWN"] reward = "score" character_set = ["@", "*", ">", "<", "^", "v"] @@ -66,13 +67,13 @@ Then create a training run, train, and watch the result: .. code-block:: console - % retro-gamer create --game my_game --output runs/snake/ + % retro-gamer init games/snake training/snake - % retro-gamer train runs/snake/ + % retro-gamer train training/snake - % retro-gamer play runs/snake/ --checkpoint ep_0500 + % retro-gamer play training/snake --checkpoint ep_0500 -The ``create`` command sets up the training run directory; ``train`` +The ``init`` command sets up the training run directory; ``train`` runs the learning algorithm; ``play`` loads a checkpoint and lets you watch the trained agent live in the terminal. diff --git a/docs/reference.rst b/docs/reference.rst index 99c3314..87d5cba 100644 --- a/docs/reference.rst +++ b/docs/reference.rst @@ -5,7 +5,7 @@ Game description fields ----------------------- Game descriptions are written in the ``[tool.retro-gamer]`` section of -your game project's ``pyproject.toml``. ``retro-gamer create`` reads +your game project's ``pyproject.toml``. ``retro-gamer init`` reads this section and copies the metadata into the training run's ``config.toml``, where it can also be inspected or hand-edited. @@ -14,6 +14,7 @@ A complete example for the Snake game: .. code-block:: toml [tool.retro-gamer] + factory = "snake:create_game" actions = ["KEY_RIGHT", "KEY_UP", "KEY_LEFT", "KEY_DOWN"] reward = "score" character_set = ["@", "*", ">", "<", "^", "v"] @@ -23,6 +24,17 @@ directly from your game's ``board_size`` attribute. The fields are described below. +``factory`` +~~~~~~~~~~~ + +**Required.** A ``"module:attr"`` string naming the function that creates +a fresh game instance. The function must take no arguments and return a +new :class:`retro.game.Game`. + +.. code-block:: toml + + factory = "snake:create_game" + ``actions`` ~~~~~~~~~~~ @@ -69,7 +81,7 @@ Preprocessing options Preprocessing options live in the ``[preprocessing]`` section of a run's ``config.toml``. They control how the game's board and state are transformed into the observation vector that the neural network sees. -``retro-gamer create`` writes sensible defaults; you can edit them by +``retro-gamer init`` writes sensible defaults; you can edit them by hand before running ``retro-gamer train``. .. note:: @@ -119,18 +131,17 @@ element each; list or tuple values are flattened. observe_state = ["apple_dx", "apple_dy"] The keys must be present in ``game.state`` at every step, initialized -in ``create_game()`` before the game starts. All values that are lists -or tuples must always have the same length from episode to episode. +before the game starts. All values that are lists or tuples must always +have the same length from episode to episode. .. warning:: - ``observe_state`` keys must be initialized to their final shape in - ``create_game()`` before the game starts. If a key is absent or its - list length changes between episodes, training will crash with an - error explaining which key changed and by how much. This happens - because the neural network's input layer has a fixed size determined - at the start of training; it cannot adapt to a changing observation - shape mid-run. + ``observe_state`` keys must be initialized to their final shape before + the game starts. If a key is absent or its list length changes between + episodes, training will crash with an error explaining which key + changed and by how much. This happens because the neural network's + input layer has a fixed size determined at the start of training; it + cannot adapt to a changing observation shape mid-run. Always initialize every observed key with a placeholder of the correct type and length before the first ``game.step()`` call. @@ -160,14 +171,20 @@ 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. +built-in encoder. -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: +``observation_function`` is the training-side complement to the game's +state dict. Use ``observe_state`` when the game already computes the +features you want and stores them in ``game.state``; use +``observation_function`` when you want to transform the game's board or +state into a representation that is more useful for learning — for example, +mapping all obstacle types to a single character, or cropping the board to +an egocentric window centred on the player. + +This is how you get an egocentric (cropped, player-centered) board — +call ``egocentric_board()`` and ``encode_board()`` yourself, from +:mod:`retro_gamer.observation`, inside your own function, and declare +``board_size`` to match your crop: .. code-block:: python @@ -194,7 +211,7 @@ own function, and declare ``board_size`` to match your crop: board_size = [17, 17] # 2*RADIUS + 1 observation_function = "my_module:egocentric_observation" -Outside DQN training — for example, BabySnake's tabular Q-learning lab, which +Outside DQN training — for example, a tabular Q-learning lab that 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 @@ -210,10 +227,11 @@ Hyperparameters are split across two sections of ``config.toml``: - ``[model]`` — network architecture (changing these requires starting fresh) - ``[training]`` — learning algorithm parameters (safe to change at any time) -Both sections can be set via ``retro-gamer create`` options or edited directly. +Both sections can be set via ``retro-gamer init`` options or edited directly +in ``config.toml`` between ``init`` and ``train``. -Learning and optimization -~~~~~~~~~~~~~~~~~~~~~~~~~ +Learning and optimization (``[training]``) +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ ``learning_rate`` (default: ``0.0001``) The step size used by the Adam optimizer when updating network @@ -232,8 +250,8 @@ Learning and optimization agent value all future rewards equally; smaller values make the agent increasingly myopic. -Exploration -~~~~~~~~~~~ +Exploration (``[training]``) +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ ``epsilon`` (default: ``1.0``) The initial exploration rate. At each turn, the agent takes a @@ -248,8 +266,8 @@ Exploration continued exploration prevents the agent from becoming permanently committed to a suboptimal policy. -Memory and sampling -~~~~~~~~~~~~~~~~~~~ +Memory and sampling (``[training]``) +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ ``batch_size`` (default: ``64``) The number of experiences sampled from the replay buffer per @@ -265,11 +283,11 @@ Memory and sampling This often improves sample efficiency at a modest computational cost. -Model architecture -~~~~~~~~~~~~~~~~~~ +Model architecture (``[model]``) +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ -These live in the ``[model]`` section. Changing them requires starting fresh -(run ``retro-gamer clean`` before retraining). +Changing any ``[model]`` option requires starting fresh (run +``retro-gamer clean`` before retraining). ``hidden_sizes`` (default: ``[128, 64]``) A list of integers giving the size of each hidden layer in the MLP @@ -278,8 +296,8 @@ These live in the ``[model]`` section. Changing them requires starting fresh network. Larger or deeper networks can represent more complex Q-functions but train more slowly and may need more episodes. -Training duration -~~~~~~~~~~~~~~~~~ +Training duration (``[training]``) +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ ``training_episodes`` (default: ``20000``) The total number of game episodes to run. Each episode runs until @@ -302,8 +320,8 @@ Training duration experience. The default of 4 is a good balance for most games; set to 1 to train on every step. -Character discovery -~~~~~~~~~~~~~~~~~~~ +Character discovery (``[training]``) +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ ``exploration_turns`` (default: ``200``) When ``character_set`` is not specified, the number of random @@ -319,8 +337,8 @@ Character discovery CLI reference ------------- -``retro-gamer create`` -~~~~~~~~~~~~~~~~~~~~~~ +``retro-gamer init`` +~~~~~~~~~~~~~~~~~~~~~ Create a new training run directory with ``config.toml``. Game metadata is read automatically from the ``[tool.retro-gamer]`` section of your @@ -328,19 +346,20 @@ game's ``pyproject.toml``; you do not pass it on the command line. .. code-block:: console - % retro-gamer create --game GAME --output DIR [OPTIONS] + % retro-gamer init GAME OUTPUT [OPTIONS] -**Required options:** +**Required arguments:** -- ``--game GAME`` — Your game, specified as a file path or a Python +- ``GAME`` — Your game, specified as a directory path or a Python module name: - - File path: ``--game my_game.py`` or ``--game my_game/`` - - Module name: ``--game retro.examples.snake`` + - Directory: ``games/snake`` + - Module name: ``retro.examples.snake`` The ``[tool.retro-gamer]`` section is read from the ``pyproject.toml`` - found in or above the game file. -- ``--output DIR`` — Directory to create for this training run. + found in or above the game directory. +- ``OUTPUT`` — Directory to create for this training run + (e.g. ``training/snake``). **Hyperparameter options** (all optional; see :ref:`hyperparameters`): @@ -369,7 +388,7 @@ Train a DQN agent. % retro-gamer train RUN_DIR ``RUN_DIR`` must contain a ``config.toml`` generated by ``retro-gamer -create``. If checkpoints already exist in ``RUN_DIR``, training +init``. If checkpoints already exist in ``RUN_DIR``, training automatically resumes from the latest one so prior work is never lost. If all configured episodes have already been completed, the command @@ -381,8 +400,8 @@ 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``, - ``observation_function``, ``extras_size`` (``[metadata]``) — game - description and observation shape + ``observation_function`` (``[metadata]``) — game description and + observation shape - ``spatial``, ``board``, ``observe_state`` (``[preprocessing]``) — observation encoding - ``hidden_sizes`` (``[model]``) — network architecture @@ -439,7 +458,7 @@ contents: .. code-block:: text - runs/snake/ + training/snake/ ├── config.toml # game description + hyperparameters ├── training.log # architecture rationale + per-episode log └── checkpoints/ @@ -447,11 +466,11 @@ contents: ├── ep_0200.pt └── ... # one file saved every 100 episodes -``config.toml`` is written by ``retro-gamer create`` and updated (with +``config.toml`` is written by ``retro-gamer init`` and updated (with the discovered character set and resolved hyperparameters) when ``retro-gamer train`` begins. It has five sections: ``[game]``, ``[metadata]``, ``[preprocessing]``, ``[model]``, and ``[training]``. -Editing ``config.toml`` between ``create`` and ``train`` is the +Editing ``config.toml`` between ``init`` and ``train`` is the recommended way to adjust hyperparameters. ``training.log`` begins with the full network architecture description, @@ -489,5 +508,5 @@ library. See the :doc:`api` reference for full details. from retro.examples.snake import create_game metadata = GameMetadata.from_pyproject("retro.examples.snake") - trainer = DQNTrainer(create_game, metadata, "runs/snake/") + trainer = DQNTrainer(create_game, metadata, "training/snake/") trainer.train() diff --git a/docs/walkthrough.rst b/docs/walkthrough.rst index 2f91592..0513c33 100644 --- a/docs/walkthrough.rst +++ b/docs/walkthrough.rst @@ -21,32 +21,29 @@ You will need: Preparing your game ------------------- -``retro-gamer`` loads your game by calling a function named -``create_game``. The function must take no arguments and return a new -``Game`` instance. +``retro-gamer`` loads your game by calling a *factory function* — a +function that takes no arguments and returns a new ``Game`` instance. +You declare this function in ``[tool.retro-gamer]``: -Here is the ``create_game`` function for Snake: +.. code-block:: toml + + [tool.retro-gamer] + factory = "snake:create_game" + +The factory function itself lives in your game's Python module: .. code-block:: python def create_game(): head = SnakeHead() apple = Apple() - game = Game([head, apple], {'score': 100}, board_size=(32, 16), framerate=12) + game = Game([head, apple], {'score': 0}, board_size=(32, 16), framerate=12) apple.relocate(game) return game -If your game file does not already have a ``create_game`` function, add -one following this pattern. - -When you run ``retro-gamer create``, you can point to your game file -directly by path or by Python module name: - -.. code-block:: console - - % retro-gamer create --game my_game.py --output runs/my_game/ - % retro-gamer create --game retro.examples.snake --output runs/snake/ - +The ``"snake:create_game"`` string follows the ``"module:attr"`` format +used throughout Python's packaging ecosystem: ``snake`` is the importable +module name and ``create_game`` is the attribute within it. Describing your game -------------------- @@ -62,12 +59,21 @@ Here is the ``[tool.retro-gamer]`` section for the Snake example: .. code-block:: toml [tool.retro-gamer] + factory = "snake:create_game" actions = ["KEY_RIGHT", "KEY_UP", "KEY_LEFT", "KEY_DOWN"] reward = "score" character_set = ["@", "*", ">", "<", "^", "v"] Let's go through each field. +``factory`` +~~~~~~~~~~~ + +The ``"module:attr"`` string naming the function that creates a fresh +game instance. ``retro-gamer`` calls this function at the start of each +training episode and whenever it needs to inspect the game (for example, +to discover the board size). + ``actions`` ~~~~~~~~~~~ @@ -129,7 +135,7 @@ The ``[tool.retro-gamer]`` section describes the game. Preprocessing options—such as ``spatial`` (whether to use a CNN or MLP, default: ``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, +``retro-gamer init``. 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. @@ -151,7 +157,7 @@ board encoding (or uses them as the entire observation when ``board = false``). These values must be set in ``game.state`` at the start of every -episode—typically inside ``create_game()``—and must keep the same +episode—typically inside the factory function—and must keep the same type and length from episode to episode. .. warning:: @@ -171,21 +177,19 @@ Once you have written this section, create the training run directory: .. code-block:: console - % retro-gamer create \ - --game retro.examples.snake \ - --output runs/snake/ + % retro-gamer init games/snake training/snake - Created training run at runs/snake/config.toml - game : retro.examples.snake + Initialized training run at training/snake/config.toml + game : games/snake board_size : 32×16 actions : ['KEY_RIGHT', 'KEY_UP', 'KEY_LEFT', 'KEY_DOWN'] reward : score characters : ['@', '*', '>', '<', '^', 'v'] architecture: MLP -``retro-gamer create`` reads your game metadata directly from +``retro-gamer init`` reads your game metadata directly from ``pyproject.toml`` and writes it—along with all hyperparameters—to -``runs/snake/config.toml``. +``training/snake/config.toml``. Training the agent ------------------ @@ -194,16 +198,16 @@ With the ``config.toml`` in place, start training: .. code-block:: console - % retro-gamer train runs/snake/ + % retro-gamer train training/snake 100%|████████████████████| 1000/1000 [12:34<00:00, 1.32ep/s, reward=9.0, eps=0.007, loss=0.0003] - Done. Checkpoints saved in runs/snake/checkpoints/ + Done. Checkpoints saved in training/snake/checkpoints/ A progress bar shows how far training has gone, along with the most recent episode's reward, the current exploration rate (``eps``), and the average prediction error (``loss``). Training saves a checkpoint every 100 episodes to -``runs/snake/checkpoints/``. You can stop training at any time with +``training/snake/checkpoints/``. You can stop training at any time with Ctrl-C and resume it later—the next ``retro-gamer train`` command will automatically pick up from the latest checkpoint. @@ -215,7 +219,7 @@ log: .. code-block:: console - % cat runs/snake/training.log + % cat training/snake/training.log The log begins with the full network architecture, followed by one line per checkpoint (every 100 episodes): @@ -264,7 +268,7 @@ checkpoint is always available immediately: .. code-block:: console - % retro-gamer play runs/snake/ + % retro-gamer play training/snake This loads the most recent checkpoint and runs the agent in your terminal. Press Enter or Escape to quit. @@ -280,7 +284,7 @@ To watch an earlier stage of training, use ``--checkpoint``: .. code-block:: console - % retro-gamer play runs/snake/ --checkpoint ep_0100 + % retro-gamer play training/snake --checkpoint ep_0100 Comparing what the agent at episode 100 does versus the agent at episode 500 can reveal exactly what the agent has (and has not) learned. For @@ -298,7 +302,7 @@ command you used before: .. code-block:: console - % retro-gamer train runs/snake/ + % retro-gamer train training/snake ``retro-gamer`` automatically detects and resumes from the latest checkpoint. No extra flags are needed. If all configured episodes have @@ -309,7 +313,7 @@ already been completed, it prints a message and exits: Training already complete (1000 episodes). To keep training, increase training_episodes in config.toml. -To continue training, open ``runs/snake/config.toml``, increase the +To continue training, open ``training/snake/config.toml``, increase the ``training_episodes`` value, and run ``retro-gamer train`` again. Watching a trained agent play @@ -319,15 +323,15 @@ Once training is complete, watch the final agent: .. code-block:: console - % retro-gamer play runs/snake/ + % retro-gamer play training/snake By default the latest checkpoint is loaded. You can also compare the agent's performance at different stages of training: .. code-block:: console - % retro-gamer play runs/snake/ --checkpoint ep_0100 - % retro-gamer play runs/snake/ --checkpoint ep_0500 + % retro-gamer play training/snake --checkpoint ep_0100 + % retro-gamer play training/snake --checkpoint ep_0500 Press Enter or Escape to quit. @@ -338,8 +342,8 @@ To review the configuration and recent training progress for a run: .. code-block:: console - % retro-gamer info runs/snake/ - Game module : retro.examples.snake + % retro-gamer info training/snake + Game module : snake Metadata : {'actions': ['KEY_RIGHT', ...], 'reward': 'score', 'board_size': [32, 16], ...} Preprocessing : {'spatial': False, 'board': True, 'observe_state': ['apple_dx', 'apple_dy'], ...} Model : {'hidden_sizes': [128, 64]} @@ -358,8 +362,8 @@ Adjusting hyperparameters -------------------------- The training hyperparameters can be changed by editing ``config.toml`` -before training, or by passing them as options to ``retro-gamer -create``. Common adjustments and their effects: +before training, or by passing them as options to ``retro-gamer init``. +Common adjustments and their effects: **``training_episodes``** — How long to train. More episodes give the agent more time to learn, but also take longer to run. This is always @@ -411,7 +415,7 @@ game or the shape of the network. The saved model weights are incompatible with the new configuration: - ``actions``, ``reward``, ``character_set``, ``board_size``, - ``observation_function``, ``extras_size`` (``[metadata]``) — These define + ``observation_function`` (``[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. @@ -443,18 +447,18 @@ To clear out the old checkpoints and begin again: .. code-block:: console - % retro-gamer clean runs/snake/ - Will remove 5 checkpoint(s) and training log from runs/snake/: + % retro-gamer clean training/snake + Will remove 5 checkpoint(s) and training log from training/snake/: checkpoints/ep_0100.pt checkpoints/ep_0200.pt ... training.log Proceed? [y/N]: y - Cleaned. Run 'retro-gamer train runs/snake/' to start fresh. + Cleaned. Run 'retro-gamer train training/snake/' to start fresh. The ``config.toml`` is always preserved so you do not need to run -``retro-gamer create`` again. +``retro-gamer init`` again. Reasoning about training from the log -------------------------------------- @@ -528,4 +532,3 @@ concepts underlying the training algorithm. episode 1000 and watch each play the same game. What has the later agent learned that the earlier one has not? How would you describe this difference to someone who does not know about neural networks? - diff --git a/pyproject.toml b/pyproject.toml index f69af0a..08cac07 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,11 +1,11 @@ [project] name = "retro-gamer" -version = "0.2.0" +version = "0.3.0" description = "A toolkit for learning reinforcement learning by training agents to play retro games" readme = "README.md" requires-python = ">=3.11" dependencies = [ - "retro-games>=2.3.1", + "retro-games>=2.5.0", "torch>=2.0", "numpy>=1.24", "click>=8.0", diff --git a/retro_gamer/cli.py b/retro_gamer/cli.py index 83b9e2a..b0e262d 100644 --- a/retro_gamer/cli.py +++ b/retro_gamer/cli.py @@ -22,15 +22,12 @@ def cli(): # --------------------------------------------------------------------------- -# retro-gamer create +# retro-gamer init # --------------------------------------------------------------------------- @cli.command() -@click.option('--game', required=True, - help='Game to train: a .py file path (e.g. my_game.py) or a Python module ' - '(e.g. retro.examples.snake)') -@click.option('--output', required=True, - help='Directory to create for this training run') +@click.argument('game') +@click.argument('output') @click.option('--learning-rate', default=DEFAULTS['learning_rate'], type=float, help=f"Adam optimizer learning rate (default {DEFAULTS['learning_rate']})") @click.option('--learning-rate-decay', default=DEFAULTS['learning_rate_decay'], type=float, @@ -60,11 +57,17 @@ def cli(): @click.option('--prioritize-experiences/--no-prioritize-experiences', default=DEFAULTS['prioritize_experiences'], help='Use prioritized experience replay') -def create(game, output, **hyperparams): +def init(game, output, **hyperparams): """Create a new training run directory. - Game metadata (actions, reward signal, etc.) is read from the - [tool.retro-gamer] section of the game's pyproject.toml. + GAME is a path to a game directory or a Python module name + (e.g. games/snake or retro.examples.snake). + + OUTPUT is the directory to create for this training run + (e.g. training/snake). + + Game metadata (actions, reward signal, factory function, etc.) is + read from the [tool.retro-gamer] section of the game's pyproject.toml. Board size is read directly from the game. Hyperparameter options control how the trainer learns, not what it learns about. """ @@ -84,6 +87,9 @@ def create(game, output, **hyperparams): except (FileNotFoundError, ValueError) as e: raise click.ClickException(str(e)) + if metadata.factory: + game_config['factory'] = metadata.factory + game_factory = _load_factory(game_config) if metadata.board_size is None: g = game_factory() @@ -105,7 +111,7 @@ def create(game, output, **hyperparams): with open(run_dir / 'config.toml', 'wb') as f: tomli_w.dump(config, f) - click.echo(f"Created training run at {output}/config.toml") + click.echo(f"Initialized training run at {output}/config.toml") click.echo(f" game : {game}") click.echo(f" board_size : {metadata.board_size[0]}×{metadata.board_size[1]}") click.echo(f" actions : {metadata.actions}") @@ -370,6 +376,21 @@ def _load_factory(game_config: dict): path = game_config.get('path') if path and path not in sys.path: sys.path.insert(0, path) + + factory_str = game_config.get('factory') + if factory_str: + module_name, attr_name = factory_str.split(':', 1) + try: + module = importlib.import_module(module_name) + except ImportError as e: + raise click.ClickException(f"Cannot import factory module '{module_name}': {e}") + if not hasattr(module, attr_name): + raise click.ClickException( + f"Module '{module_name}' has no '{attr_name}' function " + f"(from factory = {factory_str!r})" + ) + return getattr(module, attr_name) + module_name = game_config['module'] try: module = importlib.import_module(module_name) @@ -377,7 +398,8 @@ def _load_factory(game_config: dict): raise click.ClickException(f"Cannot import game module '{module_name}': {e}") if not hasattr(module, 'create_game'): raise click.ClickException( - f"Module '{module_name}' has no create_game() function" + f"Module '{module_name}' has no create_game() function. " + "Add factory = \"module:create_game\" to [tool.retro-gamer] in pyproject.toml." ) return module.create_game diff --git a/retro_gamer/metadata.py b/retro_gamer/metadata.py index 2a35def..2628460 100644 --- a/retro_gamer/metadata.py +++ b/retro_gamer/metadata.py @@ -12,15 +12,18 @@ class GameMetadata: """Describes a retro game for training purposes. Required fields: actions, reward. - Optional fields: character_set, spatial, observation_function. + Optional fields: character_set, spatial, observation_function, factory. 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). + (measured by DQNTrainer from one sampled observation — never declared). 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. + + factory, if set, is a "module:attr" string naming the create_game function. + Read from [tool.retro-gamer].factory in pyproject.toml. When absent, + the loader falls back to looking for create_game on the game module. """ actions: list[str] reward: str @@ -29,6 +32,7 @@ class GameMetadata: board: bool = True board_size: tuple[int, int] | None = None observation_function: str | None = None + factory: str | None = None extras_size: int = 0 def validate(self): @@ -76,6 +80,41 @@ class GameMetadata: "This should name a function f(game) -> observation that fully\n" "describes what your agent observes each turn." ) + if self.factory is not None: + if not isinstance(self.factory, str) or ':' not in self.factory: + raise ValueError( + f"'factory' must be a string of the form 'module:attr', " + f"but got: {self.factory!r}\n" + "Example: factory = \"my_game:create_game\"\n" + "This should name a function that takes no arguments and returns a Game." + ) + + def resolve_factory(self) -> Callable | None: + """Import and return the factory function named by the factory field, or None if unset.""" + if self.factory is None: + return None + if ':' not in self.factory: + raise ValueError( + f"'factory' must be of the form 'module:attr', but got " + f"{self.factory!r} (no ':' found).\n" + "Example: factory = \"my_game:create_game\"" + ) + module_name, attr_name = self.factory.split(':', 1) + try: + module = importlib.import_module(module_name) + except ImportError as e: + raise ValueError( + f"Could not import module {module_name!r} for factory " + f"{self.factory!r}: {e}" + ) 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 factory = {self.factory!r}).\n" + f"Define a function named '{attr_name}' in {module_name}." + ) from None def resolve_observation_function(self) -> Callable | None: """Import and return the function named by observation_function, or None if unset. @@ -153,14 +192,13 @@ class GameMetadata: spatial=d.get('spatial', False), board_size=board_size, observation_function=d.get('observation_function'), - extras_size=d.get('extras_size', 0), + factory=d.get('factory'), ) 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) diff --git a/retro_gamer/trainer.py b/retro_gamer/trainer.py index 1d52f46..1cd3a03 100644 --- a/retro_gamer/trainer.py +++ b/retro_gamer/trainer.py @@ -55,7 +55,6 @@ _INCOMPATIBLE_METADATA = { '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)', diff --git a/uv.lock b/uv.lock index 8a4cc0f..0b21275 100644 --- a/uv.lock +++ b/uv.lock @@ -1154,7 +1154,7 @@ wheels = [ [[package]] name = "retro-gamer" -version = "0.2.0" +version = "0.3.0" source = { editable = "." } dependencies = [ { name = "click" }, @@ -1196,7 +1196,7 @@ documentation = [ [[package]] name = "retro-games" -version = "2.4.1" +version = "2.5.0" source = { editable = "../retro" } dependencies = [ { name = "blessed" }, @@ -1207,8 +1207,8 @@ requires-dist = [{ name = "blessed", specifier = ">=1.33.0" }] [package.metadata.requires-dev] documentation = [ - { name = "furo", specifier = ">=2025.12.19" }, { name = "sphinx", specifier = ">=8.1.3" }, + { name = "sphinx-rtd-theme", specifier = ">=3.0" }, ] [[package]]