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