Updates
This commit is contained in:
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user