This commit is contained in:
Chris Proctor
2026-06-26 20:59:25 -04:00
parent a9385f8296
commit 482f4f6cfa
5 changed files with 191 additions and 86 deletions

View File

@@ -2,9 +2,82 @@ import numpy as np
from retro.views.headless import HeadlessView
from retro_gamer.observation import egocentric_board, encode_board, encode_state
HEAD_CHARS = frozenset({">", "<", "^", "v"})
CHARACTER_SET = ["@", "*", ">", "<", "^", "v"]
CHARACTER_SET_NORMALIZED = ["@", "*", "H"]
RADIUS = 8
DIRECTIONS = [(1, 0), (0, -1), (-1, 0), (0, 1)] # RIGHT, UP, LEFT, DOWN
RADIUS_WIDE = 8 # 17×17 egocentric window
RADIUS_NARROW = 3 # 7×7 egocentric window
def direction_observation(game):
"""Normalized 3-char board plus apple_dx, apple_dy, and a one-hot direction.
Returns board in channel-first (C, H, W) order for the CNN, followed by
apple_dx, apple_dy, and 4 one-hot bits encoding the snake's heading.
The board collapses all four head characters to 'H', keeping 3 channels
instead of 6 while the one-hot direction restores that information cheaply.
"""
view = HeadlessView()
view.on_game_start(game)
view.render(game)
normalized = [
["H" if c in HEAD_CHARS else c for c in row]
for row in view.board_characters
]
board_vec = encode_board(normalized, CHARACTER_SET_NORMALIZED).transpose(2, 0, 1).flatten()
extras = encode_state(game.state, ["apple_dx", "apple_dy"])
head = game.get_agent_by_name("Snake head")
direction_onehot = np.array(
[1.0 if head.direction == d else 0.0 for d in DIRECTIONS],
dtype=np.float32,
)
return np.concatenate([board_vec, extras, direction_onehot])
def normalized_observation(game):
"""Full board with all head chars collapsed to 'H', plus apple_dx and apple_dy.
Returns board in channel-first (C, H, W) order for the CNN, followed by extras.
Character set: ['@' apple, '*' body, 'H' head] — 3 channels instead of 6.
"""
view = HeadlessView()
view.on_game_start(game)
view.render(game)
normalized = [
["H" if c in HEAD_CHARS else c for c in row]
for row in view.board_characters
]
board_vec = encode_board(normalized, CHARACTER_SET_NORMALIZED).transpose(2, 0, 1).flatten()
extras = encode_state(game.state, ["apple_dx", "apple_dy"])
return np.concatenate([board_vec, extras])
def egocentric_cnn_observation(game):
"""17×17 egocentric window in channel-first (CHW) format for CNN, plus apple_dx/dy."""
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_WIDE)
board_vec = encode_board(cropped, CHARACTER_SET).transpose(2, 0, 1).flatten()
extras = encode_state(game.state, ["apple_dx", "apple_dy"])
return np.concatenate([board_vec, extras])
def narrow_egocentric_observation(game):
"""7×7 egocentric window (flat, for MLP), plus apple_dx and apple_dy."""
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_NARROW)
board_vec = encode_board(cropped, CHARACTER_SET).flatten()
extras = encode_state(game.state, ["apple_dx", "apple_dy"])
return np.concatenate([board_vec, extras])
def egocentric_observation(game):
"""17×17 window centered on the snake's head, plus apple_dx and apple_dy."""

View File

@@ -8,3 +8,4 @@ dependencies = ["retro-games>=2.5.0"]
[tool.retro-gamer]
actions = ["KEY_RIGHT", "KEY_UP", "KEY_LEFT", "KEY_DOWN"]
reward = "reward"
character_set = ["@", "*", "<", ">", "^", "v"]