Team Ai
Apppublic

openenv/atari_env

sourceHugging Faceupdated 6mo agoView on Hugging Face
3likes
atari_environment.py255 linesDownload Raw Back to server
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