Files
retro-gamer/retro_gamer/metadata.py
Chris Proctor c89609fe77 Refactor CLI/interface: init command, factory field, remove extras_size
- 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
2026-06-26 06:55:30 -04:00

246 lines
11 KiB
Python

from __future__ import annotations
import importlib
import tomllib
import tomli_w
from dataclasses import dataclass, field
from pathlib import Path
from typing import Callable
@dataclass
class GameMetadata:
"""Describes a retro game for training purposes.
Required fields: actions, reward.
Optional fields: character_set, spatial, observation_function, factory.
Discovered fields: board_size (from game.board_size), extras_size
(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
character_set: list[str] | None = None
spatial: bool = False
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):
if not self.actions:
raise ValueError(
"The 'actions' list in [tool.retro-gamer] is empty or missing.\n"
"It should list the keyboard keys your agent can press, for example:\n\n"
' actions = ["KEY_RIGHT", "KEY_UP", "KEY_LEFT", "KEY_DOWN"]\n\n'
"The agent will learn which actions lead to higher rewards."
)
if not isinstance(self.actions, list) or not all(isinstance(a, str) for a in self.actions):
raise ValueError(
f"'actions' must be a list of strings, but got: {self.actions!r}\n"
"Each entry should be a key name like \"KEY_RIGHT\" or \"KEY_SPACE\"."
)
if not self.reward:
raise ValueError(
"The 'reward' field in [tool.retro-gamer] is empty or missing.\n"
"It should name a game state variable whose value the agent is trying\n"
"to maximize — for example:\n\n"
" reward = \"score\"\n\n"
"The trainer watches how this value changes each step and uses those\n"
"changes as the reward signal."
)
if self.character_set is not None:
if not isinstance(self.character_set, list):
raise ValueError(
f"'character_set' must be a list of single characters, but got: {self.character_set!r}\n"
"Example: character_set = [\"@\", \"*\", \"#\"]"
)
for ch in self.character_set:
if not isinstance(ch, str) or len(ch) != 1:
raise ValueError(
f"Every entry in character_set must be a single character, but got {ch!r}.\n"
"Each character represents one type of cell on the game board.\n"
"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."
)
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.
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:
"""Load metadata from the [tool.retro-gamer] section of the game's pyproject.toml."""
pyproject_path = _find_pyproject(module_name)
if pyproject_path is None:
raise FileNotFoundError(
f"Could not find pyproject.toml for module '{module_name}'. "
f"Make sure the module is part of a Python project with a pyproject.toml."
)
with open(pyproject_path, 'rb') as f:
data = tomllib.load(f)
section = data.get('tool', {}).get('retro-gamer')
if section is None:
raise ValueError(
f"No [tool.retro-gamer] section found in {pyproject_path}.\n"
f"Add game metadata to your pyproject.toml:\n\n"
f"[tool.retro-gamer]\n"
f"actions = [\"KEY_RIGHT\", ...]\n"
f"reward = \"score\"\n"
)
return cls.from_dict(section)
@classmethod
def from_dict(cls, d: dict) -> GameMetadata:
missing = [k for k in ('actions', 'reward') if k not in d]
if missing:
fields = ' and '.join(f"'{k}'" for k in missing)
raise ValueError(
f"The [tool.retro-gamer] section is missing required {fields}.\n"
"A minimal configuration looks like this:\n\n"
"[tool.retro-gamer]\n"
'actions = ["KEY_RIGHT", "KEY_UP", "KEY_LEFT", "KEY_DOWN"]\n'
'reward = "score"\n\n'
"See the documentation for all available options."
)
board_size = tuple(d['board_size']) if 'board_size' in d else None
return cls(
actions=d['actions'],
reward=d['reward'],
character_set=d.get('character_set'),
spatial=d.get('spatial', False),
board_size=board_size,
observation_function=d.get('observation_function'),
factory=d.get('factory'),
)
def to_dict(self) -> dict:
d = {
'actions': self.actions,
'reward': self.reward,
}
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):
with open(path, 'wb') as f:
tomli_w.dump({'metadata': self.to_dict()}, f)
@property
def obs_size(self) -> int:
"""Total size of the flat observation vector."""
if not self.board:
return self.extras_size
C = len(self.character_set) if self.character_set else 0
bw, bh = self.board_size
return C * bw * bh + self.extras_size
@property
def n_actions(self) -> int:
"""Number of actions including no-op."""
return len(self.actions) + 1
def _find_pyproject(module_name: str) -> Path | None:
"""Walk up from a module's source file to find its pyproject.toml."""
try:
module = importlib.import_module(module_name)
except ImportError:
return None
module_file = getattr(module, '__file__', None)
if module_file is None:
return None
for parent in Path(module_file).resolve().parents:
candidate = parent / 'pyproject.toml'
if candidate.exists():
return candidate
if parent.name == 'site-packages':
break
return None