openenv/atari_env
3
1# Copyright (c) Meta Platforms, Inc. and affiliates.2# All rights reserved.3#4# This source code is licensed under the BSD-style license found in the5# LICENSE file in the root directory of this source tree.6 7"""8Atari Environment Server Implementation.9 10This module wraps ALE's ALEInterface and exposes it11via the OpenEnv Environment interface.12"""13 14import uuid15from typing import Any, Dict, Literal, Optional16 17from openenv.core.env_server import Action, Environment, Observation18 19# Support both in-repo and standalone imports20try:21 # In-repo imports (when running from OpenEnv repository)22 from ..models import AtariAction, AtariObservation, AtariState23except ImportError as e:24 if "relative import" not in str(e) and "no known parent package" not in str(e):25 raise26 # Standalone imports (when running via uvicorn server.app:app)27 from models import AtariAction, AtariObservation, AtariState28 29# Import ALE30try:31 import numpy as np32 from ale_py import ALEInterface, roms33except ImportError as e:34 raise ImportError(35 "ALE (Arcade Learning Environment) is not installed. "36 "Please install it with: pip install ale-py"37 ) from e38 39 40class AtariEnvironment(Environment):41 """42 Atari Environment wrapper for OpenEnv.43 44 This environment wraps Atari 2600 games via the Arcade Learning Environment (ALE)45 and provides a clean interface for RL training.46 47 Supported games include: pong, breakout, space_invaders, and 100+ others.48 49 Args:50 game_name: Name of the Atari game (e.g., "pong", "breakout").51 obs_type: Observation type - "rgb", "grayscale", or "ram".52 full_action_space: Use full action space (18 actions) vs minimal.53 mode: Game mode (if applicable).54 difficulty: Game difficulty (if applicable).55 repeat_action_probability: Sticky action probability (default 0.0).56 frameskip: Number of frames to skip per action (default 4).57 58 Example:59 >>> env = AtariEnvironment("pong")60 >>> obs = env.reset()61 >>> print(obs.screen_shape) # [210, 160, 3]62 >>> obs = env.step(AtariAction(action_id=2)) # UP63 >>> print(obs.reward, obs.done)64 """65 66 def __init__(67 self,68 game_name: str = "pong",69 obs_type: Literal["rgb", "grayscale", "ram"] = "rgb",70 full_action_space: bool = False,71 mode: Optional[int] = None,72 difficulty: Optional[int] = None,73 repeat_action_probability: float = 0.0,74 frameskip: int = 4,75 ):76 """Initialize Atari environment."""77 super().__init__()78 79 self.game_name = game_name80 self.obs_type = obs_type81 self.full_action_space = full_action_space82 self.mode = mode83 self.difficulty = difficulty84 self.repeat_action_probability = repeat_action_probability85 self.frameskip = frameskip86 87 # Create ALE interface88 self.ale = ALEInterface()89 90 # Configure ALE91 from ale_py import LoggerMode92 93 self.ale.setLoggerMode(LoggerMode.Error) # Error mode only94 self.ale.setFloat("repeat_action_probability", repeat_action_probability)95 96 # Load ROM97 try:98 rom_path = roms.get_rom_path(game_name)99 self.ale.loadROM(rom_path)100 except Exception as e:101 raise ValueError(102 f"Failed to load Atari game '{game_name}': {e}\n"103 f"Available games can be found via: ale_py.roms.list_roms()"104 ) from e105 106 # Set mode and difficulty if specified107 if mode is not None:108 self.ale.setMode(mode)109 if difficulty is not None:110 self.ale.setDifficulty(difficulty)111 112 # Get action set113 if full_action_space:114 self._action_set = self.ale.getLegalActionSet()115 else:116 self._action_set = self.ale.getMinimalActionSet()117 118 # Get screen dimensions for observation space119 self.screen_height, self.screen_width = self.ale.getScreenDims()120 if obs_type == "rgb":121 self.screen_shape = [self.screen_height, self.screen_width, 3]122 elif obs_type == "grayscale":123 self.screen_shape = [self.screen_height, self.screen_width]124 elif obs_type == "ram":125 self.screen_shape = [self.ale.getRAMSize()]126 else:127 raise ValueError(f"Invalid obs_type: {obs_type}")128 129 # Initialize state130 self._state = AtariState(131 game_name=game_name,132 obs_type=obs_type,133 full_action_space=full_action_space,134 mode=mode,135 difficulty=difficulty,136 repeat_action_probability=repeat_action_probability,137 frameskip=frameskip,138 )139 140 def reset(self) -> Observation:141 """142 Reset the environment and return initial observation.143 144 Returns:145 Initial observation for the agent.146 """147 # Reset ALE148 self.ale.reset_game()149 150 # Reset state tracking151 self._state.episode_id = str(uuid.uuid4())152 self._state.step_count = 0153 154 # Get initial observation155 return self._make_observation()156 157 def step(self, action: Action) -> Observation:158 """159 Execute agent's action and return resulting observation.160 161 Args:162 action: AtariAction containing the action_id to execute.163 164 Returns:165 Observation after action execution.166 167 Raises:168 ValueError: If action is not an AtariAction.169 """170 if not isinstance(action, AtariAction):171 raise ValueError(f"Expected AtariAction, got {type(action)}")172 173 # Validate action_id174 if action.action_id < 0 or action.action_id >= len(self._action_set):175 raise ValueError(176 f"Invalid action_id: {action.action_id}. "177 f"Valid range: [0, {len(self._action_set) - 1}]"178 )179 180 # Get actual ALE action181 ale_action = self._action_set[action.action_id]182 183 # Execute action with frameskip184 total_reward = 0.0185 for _ in range(self.frameskip):186 total_reward += self.ale.act(ale_action)187 if self.ale.game_over():188 break189 190 self._state.step_count += 1191 192 # Get observation193 obs = self._make_observation()194 obs.reward = total_reward195 196 return obs197 198 @property199 def state(self) -> AtariState:200 """Get current environment state."""201 return self._state202 203 def _make_observation(self) -> AtariObservation:204 """205 Create an AtariObservation from current ALE state.206 207 Returns:208 AtariObservation for the agent.209 """210 # Get screen observation211 if self.obs_type == "rgb":212 screen = self.ale.getScreenRGB()213 elif self.obs_type == "grayscale":214 screen = self.ale.getScreenGrayscale()215 elif self.obs_type == "ram":216 screen = self.ale.getRAM()217 else:218 raise ValueError(f"Invalid obs_type: {self.obs_type}")219 220 # Flatten screen for JSON serialization221 # Handle both numpy arrays and lists222 if hasattr(screen, "flatten"):223 screen_flat = screen.flatten().tolist()224 elif hasattr(screen, "tolist"):225 screen_flat = screen.tolist()226 else:227 screen_flat = list(screen)228 229 # Get game info230 lives = self.ale.lives()231 episode_frame_number = self.ale.getEpisodeFrameNumber()232 frame_number = self.ale.getFrameNumber()233 done = self.ale.game_over()234 235 # Create legal actions list (indices into action_set)236 legal_actions = list(range(len(self._action_set)))237 238 # Create observation239 obs = AtariObservation(240 screen=screen_flat,241 screen_shape=self.screen_shape,242 legal_actions=legal_actions,243 lives=lives,244 episode_frame_number=episode_frame_number,245 frame_number=frame_number,246 done=done,247 reward=0.0, # Will be filled in by step()248 metadata={249 "game_name": self.game_name,250 "action_meanings": [str(a) for a in self._action_set],251 },252 )253 254 return obs255 