From 470771b5c2142f0813b6bb1afb9143c960bf09a0 Mon Sep 17 00:00:00 2001 From: Antonin RAFFIN Date: Sun, 12 Mar 2023 18:47:52 +0100 Subject: [PATCH] Fix Atari Roms download, enable RUF linting (#1379) * Add extra no Atari and fix CI for forks * Enable ruff rules * Change to no roms --- .github/workflows/ci.yml | 5 +-- docs/misc/changelog.rst | 3 +- pyproject.toml | 4 +- setup.py | 40 ++++++++++++------- stable_baselines3/common/atari_wrappers.py | 2 +- stable_baselines3/common/buffers.py | 20 +++++----- stable_baselines3/common/callbacks.py | 2 +- stable_baselines3/common/monitor.py | 2 +- stable_baselines3/common/policies.py | 8 ++-- stable_baselines3/common/results_plotter.py | 2 +- stable_baselines3/common/save_util.py | 2 +- .../common/vec_env/dummy_vec_env.py | 2 +- .../common/vec_env/stacked_observations.py | 2 +- stable_baselines3/dqn/dqn.py | 2 +- stable_baselines3/her/her_replay_buffer.py | 12 +++--- stable_baselines3/sac/sac.py | 2 +- stable_baselines3/td3/td3.py | 2 +- stable_baselines3/version.txt | 2 +- tests/test_logger.py | 2 +- tests/test_utils.py | 8 ++-- tests/test_vec_envs.py | 4 +- 21 files changed, 69 insertions(+), 59 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index fc47e9d..52f5f38 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -14,7 +14,6 @@ jobs: env: TERM: xterm-256color FORCE_COLOR: 1 - ATARI_ROMS: ${{ secrets.ATARI_ROMS }} # Skip CI if [ci skip] in the commit message if: "! contains(toJSON(github.event.commits.*.message), '[ci skip]')" @@ -37,11 +36,11 @@ jobs: # Install Atari Roms pip install autorom - wget $ATARI_ROMS + wget https://gist.githubusercontent.com/jjshoots/61b22aefce4456920ba99f2c36906eda/raw/00046ac3403768bfe45857610a3d333b8e35e026/Roms.tar.gz.b64 base64 Roms.tar.gz.b64 --decode &> Roms.tar.gz AutoROM --accept-license --source-file Roms.tar.gz - pip install .[extra,tests,docs] + pip install .[extra_no_roms,tests,docs] # Use headless version pip install opencv-python-headless - name: Lint with ruff diff --git a/docs/misc/changelog.rst b/docs/misc/changelog.rst index c6f8211..9573363 100644 --- a/docs/misc/changelog.rst +++ b/docs/misc/changelog.rst @@ -4,7 +4,7 @@ Changelog ========== -Release 1.8.0a8 (WIP) +Release 1.8.0a9 (WIP) -------------------------- @@ -46,6 +46,7 @@ Others: - Moved from ``setup.cg`` to ``pyproject.toml`` configuration file - Switched from ``flake8`` to ``ruff`` - Upgraded AutoROM to latest version +- Added ``extra_no_roms`` option for package installation without Atari Roms Documentation: ^^^^^^^^^^^^^^ diff --git a/pyproject.toml b/pyproject.toml index 461e727..67a7edc 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -3,8 +3,8 @@ line-length = 127 # Assume Python 3.7 target-version = "py37" -# TODO(antonin): activate "RUF" https://beta.ruff.rs/docs/rules/#ruff-specific-rules-ruf -select = ["E", "F", "B", "UP", "C90"] +# See https://beta.ruff.rs/docs/rules/ +select = ["E", "F", "B", "UP", "C90", "RUF"] ignore = [] [tool.ruff.per-file-ignores] diff --git a/setup.py b/setup.py index dcff0a2..7e00433 100644 --- a/setup.py +++ b/setup.py @@ -70,6 +70,29 @@ model = PPO("MlpPolicy", "CartPole-v1").learn(10_000) """ # noqa:E501 +# Atari Games download is sometimes problematic: +# https://github.com/Farama-Foundation/AutoROM/issues/39 +# That's why we define extra packages without it. +extra_no_roms = [ + # For render + "opencv-python", + # Tensorboard support + "tensorboard>=2.9.1", + # Checking memory taken by replay buffer + "psutil", + # For progress bar callback + "tqdm", + "rich", + # For atari games, + "ale-py==0.7.4", + "pillow", +] + +extra_packages = extra_no_roms + [ # noqa: RUF005 + # For atari roms, + "autorom[accept-rom-license]~=0.5.5", +] + setup( name="stable_baselines3", @@ -119,21 +142,8 @@ setup( # Copy button for code snippets "sphinx_copybutton", ], - "extra": [ - # For render - "opencv-python", - # For atari games, - "ale-py==0.7.4", - "autorom[accept-rom-license]~=0.5.5", - "pillow", - # Tensorboard support - "tensorboard>=2.9.1", - # Checking memory taken by replay buffer - "psutil", - # For progress bar callback - "tqdm", - "rich", - ], + "extra": extra_packages, + "extra_no_roms": extra_no_roms, }, description="Pytorch version of Stable Baselines, implementations of reinforcement learning algorithms.", author="Antonin Raffin", diff --git a/stable_baselines3/common/atari_wrappers.py b/stable_baselines3/common/atari_wrappers.py index 1e06a7f..ad29a31 100644 --- a/stable_baselines3/common/atari_wrappers.py +++ b/stable_baselines3/common/atari_wrappers.py @@ -156,7 +156,7 @@ class MaxAndSkipEnv(gym.Wrapper): def __init__(self, env: gym.Env, skip: int = 4) -> None: super().__init__(env) # most recent raw observations (for max pooling across time steps) - self._obs_buffer = np.zeros((2,) + env.observation_space.shape, dtype=env.observation_space.dtype) + self._obs_buffer = np.zeros((2, *env.observation_space.shape), dtype=env.observation_space.dtype) self._skip = skip def step(self, action: int) -> GymStepReturn: diff --git a/stable_baselines3/common/buffers.py b/stable_baselines3/common/buffers.py index 273dba9..f14d2f9 100644 --- a/stable_baselines3/common/buffers.py +++ b/stable_baselines3/common/buffers.py @@ -67,7 +67,7 @@ class BaseBuffer(ABC): """ shape = arr.shape if len(shape) < 3: - shape = shape + (1,) + shape = (*shape, 1) return arr.swapaxes(0, 1).reshape(shape[0] * shape[1], *shape[2:]) def size(self) -> int: @@ -199,13 +199,13 @@ class ReplayBuffer(BaseBuffer): ) self.optimize_memory_usage = optimize_memory_usage - self.observations = np.zeros((self.buffer_size, self.n_envs) + self.obs_shape, dtype=observation_space.dtype) + self.observations = np.zeros((self.buffer_size, self.n_envs, *self.obs_shape), dtype=observation_space.dtype) if optimize_memory_usage: # `observations` contains also the next observation self.next_observations = None else: - self.next_observations = np.zeros((self.buffer_size, self.n_envs) + self.obs_shape, dtype=observation_space.dtype) + self.next_observations = np.zeros((self.buffer_size, self.n_envs, *self.obs_shape), dtype=observation_space.dtype) self.actions = np.zeros((self.buffer_size, self.n_envs, self.action_dim), dtype=action_space.dtype) @@ -243,8 +243,8 @@ class ReplayBuffer(BaseBuffer): # Reshape needed when using multiple envs with discrete observations # as numpy cannot broadcast (n_discrete,) to (n_discrete, 1) if isinstance(self.observation_space, spaces.Discrete): - obs = obs.reshape((self.n_envs,) + self.obs_shape) - next_obs = next_obs.reshape((self.n_envs,) + self.obs_shape) + obs = obs.reshape((self.n_envs, *self.obs_shape)) + next_obs = next_obs.reshape((self.n_envs, *self.obs_shape)) # Same, for actions action = action.reshape((self.n_envs, self.action_dim)) @@ -354,7 +354,7 @@ class RolloutBuffer(BaseBuffer): self.reset() def reset(self) -> None: - self.observations = np.zeros((self.buffer_size, self.n_envs) + self.obs_shape, dtype=np.float32) + self.observations = np.zeros((self.buffer_size, self.n_envs, *self.obs_shape), dtype=np.float32) self.actions = np.zeros((self.buffer_size, self.n_envs, self.action_dim), dtype=np.float32) self.rewards = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32) self.returns = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32) @@ -428,7 +428,7 @@ class RolloutBuffer(BaseBuffer): # Reshape needed when using multiple envs with discrete observations # as numpy cannot broadcast (n_discrete,) to (n_discrete, 1) if isinstance(self.observation_space, spaces.Discrete): - obs = obs.reshape((self.n_envs,) + self.obs_shape) + obs = obs.reshape((self.n_envs, *self.obs_shape)) # Same reshape, for actions action = action.reshape((self.n_envs, self.action_dim)) @@ -528,11 +528,11 @@ class DictReplayBuffer(ReplayBuffer): self.optimize_memory_usage = optimize_memory_usage self.observations = { - key: np.zeros((self.buffer_size, self.n_envs) + _obs_shape, dtype=observation_space[key].dtype) + key: np.zeros((self.buffer_size, self.n_envs, *_obs_shape), dtype=observation_space[key].dtype) for key, _obs_shape in self.obs_shape.items() } self.next_observations = { - key: np.zeros((self.buffer_size, self.n_envs) + _obs_shape, dtype=observation_space[key].dtype) + key: np.zeros((self.buffer_size, self.n_envs, *_obs_shape), dtype=observation_space[key].dtype) for key, _obs_shape in self.obs_shape.items() } @@ -699,7 +699,7 @@ class DictRolloutBuffer(RolloutBuffer): assert isinstance(self.obs_shape, dict), "DictRolloutBuffer must be used with Dict obs space only" self.observations = {} for key, obs_input_shape in self.obs_shape.items(): - self.observations[key] = np.zeros((self.buffer_size, self.n_envs) + obs_input_shape, dtype=np.float32) + self.observations[key] = np.zeros((self.buffer_size, self.n_envs, *obs_input_shape), dtype=np.float32) self.actions = np.zeros((self.buffer_size, self.n_envs, self.action_dim), dtype=np.float32) self.rewards = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32) self.returns = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32) diff --git a/stable_baselines3/common/callbacks.py b/stable_baselines3/common/callbacks.py index 69a21ab..28d07e0 100644 --- a/stable_baselines3/common/callbacks.py +++ b/stable_baselines3/common/callbacks.py @@ -554,7 +554,7 @@ class StopTrainingOnRewardThreshold(BaseCallback): class EveryNTimesteps(EventCallback): """ - Trigger a callback every ``n_steps`` timesteps + Trigger a callback every ``n_steps`` timesteps :param n_steps: Number of timesteps between two trigger. :param callback: Callback that will be called diff --git a/stable_baselines3/common/monitor.py b/stable_baselines3/common/monitor.py index ffc318d..b8ebc2b 100644 --- a/stable_baselines3/common/monitor.py +++ b/stable_baselines3/common/monitor.py @@ -194,7 +194,7 @@ class ResultsWriter: mode = "w" if override_existing else "a" # Prevent newline issue on Windows, see GH issue #692 self.file_handler = open(filename, f"{mode}t", newline="\n") - self.logger = csv.DictWriter(self.file_handler, fieldnames=("r", "l", "t") + extra_keys) + self.logger = csv.DictWriter(self.file_handler, fieldnames=("r", "l", "t", *extra_keys)) if override_existing: self.file_handler.write(f"#{json.dumps(header)}\n") self.logger.writeheader() diff --git a/stable_baselines3/common/policies.py b/stable_baselines3/common/policies.py index 795ed47..e07aefa 100644 --- a/stable_baselines3/common/policies.py +++ b/stable_baselines3/common/policies.py @@ -228,7 +228,7 @@ class BaseModel(nn.Module): obs_ = np.array(obs) vectorized_env = vectorized_env or is_vectorized_observation(obs_, obs_space) # Add batch dimension if needed - observation[key] = obs_.reshape((-1,) + self.observation_space[key].shape) + observation[key] = obs_.reshape((-1, *self.observation_space[key].shape)) elif is_image_space(self.observation_space): # Handle the different cases for images @@ -242,7 +242,7 @@ class BaseModel(nn.Module): # Dict obs need to be handled separately vectorized_env = is_vectorized_observation(observation, self.observation_space) # Add batch dimension if needed - observation = observation.reshape((-1,) + self.observation_space.shape) + observation = observation.reshape((-1, *self.observation_space.shape)) observation = obs_as_tensor(observation, self.device) return observation, vectorized_env @@ -330,7 +330,7 @@ class BasePolicy(BaseModel, ABC): with th.no_grad(): actions = self._predict(observation, deterministic=deterministic) # Convert to numpy, and reshape to the original action shape - actions = actions.cpu().numpy().reshape((-1,) + self.action_space.shape) + actions = actions.cpu().numpy().reshape((-1, *self.action_space.shape)) if isinstance(self.action_space, spaces.Box): if self.squash_output: @@ -608,7 +608,7 @@ class ActorCriticPolicy(BasePolicy): distribution = self._get_action_dist_from_latent(latent_pi) actions = distribution.get_actions(deterministic=deterministic) log_prob = distribution.log_prob(actions) - actions = actions.reshape((-1,) + self.action_space.shape) + actions = actions.reshape((-1, *self.action_space.shape)) return actions, values, log_prob def extract_features(self, obs: th.Tensor) -> Union[th.Tensor, Tuple[th.Tensor, th.Tensor]]: diff --git a/stable_baselines3/common/results_plotter.py b/stable_baselines3/common/results_plotter.py index dac2b6c..1324557 100644 --- a/stable_baselines3/common/results_plotter.py +++ b/stable_baselines3/common/results_plotter.py @@ -25,7 +25,7 @@ def rolling_window(array: np.ndarray, window: int) -> np.ndarray: :return: rolling window on the input array """ shape = array.shape[:-1] + (array.shape[-1] - window + 1, window) - strides = array.strides + (array.strides[-1],) + strides = (*array.strides, array.strides[-1]) return np.lib.stride_tricks.as_strided(array, shape=shape, strides=strides) diff --git a/stable_baselines3/common/save_util.py b/stable_baselines3/common/save_util.py index 7ae1e22..e5aeb66 100644 --- a/stable_baselines3/common/save_util.py +++ b/stable_baselines3/common/save_util.py @@ -37,7 +37,7 @@ def recursive_getattr(obj: Any, attr: str, *args) -> Any: def _getattr(obj: Any, attr: str) -> Any: return getattr(obj, attr, *args) - return functools.reduce(_getattr, [obj] + attr.split(".")) + return functools.reduce(_getattr, [obj, *attr.split(".")]) def recursive_setattr(obj: Any, attr: str, val: Any) -> None: diff --git a/stable_baselines3/common/vec_env/dummy_vec_env.py b/stable_baselines3/common/vec_env/dummy_vec_env.py index 72cab7e..5b9fc8b 100644 --- a/stable_baselines3/common/vec_env/dummy_vec_env.py +++ b/stable_baselines3/common/vec_env/dummy_vec_env.py @@ -39,7 +39,7 @@ class DummyVecEnv(VecEnv): obs_space = env.observation_space self.keys, shapes, dtypes = obs_space_info(obs_space) - self.buf_obs = OrderedDict([(k, np.zeros((self.num_envs,) + tuple(shapes[k]), dtype=dtypes[k])) for k in self.keys]) + self.buf_obs = OrderedDict([(k, np.zeros((self.num_envs, *tuple(shapes[k])), dtype=dtypes[k])) for k in self.keys]) self.buf_dones = np.zeros((self.num_envs,), dtype=bool) self.buf_rews = np.zeros((self.num_envs,), dtype=np.float32) self.buf_infos = [{} for _ in range(self.num_envs)] diff --git a/stable_baselines3/common/vec_env/stacked_observations.py b/stable_baselines3/common/vec_env/stacked_observations.py index a26812c..ae7aebc 100644 --- a/stable_baselines3/common/vec_env/stacked_observations.py +++ b/stable_baselines3/common/vec_env/stacked_observations.py @@ -56,7 +56,7 @@ class StackedObservations(Generic[TObs]): low = np.repeat(observation_space.low, n_stack, axis=self.repeat_axis) high = np.repeat(observation_space.high, n_stack, axis=self.repeat_axis) self.stacked_observation_space = spaces.Box(low=low, high=high, dtype=observation_space.dtype) - self.stacked_obs = np.zeros((num_envs,) + self.stacked_shape, dtype=observation_space.dtype) + self.stacked_obs = np.zeros((num_envs, *self.stacked_shape), dtype=observation_space.dtype) else: raise TypeError( f"StackedObservations only supports Box and Dict as observation spaces. {observation_space} was provided." diff --git a/stable_baselines3/dqn/dqn.py b/stable_baselines3/dqn/dqn.py index ea1946a..61560e2 100644 --- a/stable_baselines3/dqn/dqn.py +++ b/stable_baselines3/dqn/dqn.py @@ -270,7 +270,7 @@ class DQN(OffPolicyAlgorithm): ) def _excluded_save_params(self) -> List[str]: - return super()._excluded_save_params() + ["q_net", "q_net_target"] + return [*super()._excluded_save_params(), "q_net", "q_net_target"] def _get_torch_save_params(self) -> Tuple[List[str], List[str]]: state_dicts = ["policy", "policy.optimizer"] diff --git a/stable_baselines3/her/her_replay_buffer.py b/stable_baselines3/her/her_replay_buffer.py index 0c3da25..d4d1a43 100644 --- a/stable_baselines3/her/her_replay_buffer.py +++ b/stable_baselines3/her/her_replay_buffer.py @@ -130,14 +130,14 @@ class HerReplayBuffer(DictReplayBuffer): # input dimensions for buffer initialization input_shape = { - "observation": (self.env.num_envs,) + self.obs_shape, - "achieved_goal": (self.env.num_envs,) + self.goal_shape, - "desired_goal": (self.env.num_envs,) + self.goal_shape, + "observation": (self.env.num_envs, *self.obs_shape), + "achieved_goal": (self.env.num_envs, *self.goal_shape), + "desired_goal": (self.env.num_envs, *self.goal_shape), "action": (self.action_dim,), "reward": (1,), - "next_obs": (self.env.num_envs,) + self.obs_shape, - "next_achieved_goal": (self.env.num_envs,) + self.goal_shape, - "next_desired_goal": (self.env.num_envs,) + self.goal_shape, + "next_obs": (self.env.num_envs, *self.obs_shape), + "next_achieved_goal": (self.env.num_envs, *self.goal_shape), + "next_desired_goal": (self.env.num_envs, *self.goal_shape), "done": (1,), } self._observation_keys = ["observation", "achieved_goal", "desired_goal"] diff --git a/stable_baselines3/sac/sac.py b/stable_baselines3/sac/sac.py index d1a6610..85a4584 100644 --- a/stable_baselines3/sac/sac.py +++ b/stable_baselines3/sac/sac.py @@ -304,7 +304,7 @@ class SAC(OffPolicyAlgorithm): ) def _excluded_save_params(self) -> List[str]: - return super()._excluded_save_params() + ["actor", "critic", "critic_target"] + return super()._excluded_save_params() + ["actor", "critic", "critic_target"] # noqa: RUF005 def _get_torch_save_params(self) -> Tuple[List[str], List[str]]: state_dicts = ["policy", "actor.optimizer", "critic.optimizer"] diff --git a/stable_baselines3/td3/td3.py b/stable_baselines3/td3/td3.py index c844a99..30ae867 100644 --- a/stable_baselines3/td3/td3.py +++ b/stable_baselines3/td3/td3.py @@ -218,7 +218,7 @@ class TD3(OffPolicyAlgorithm): ) def _excluded_save_params(self) -> List[str]: - return super()._excluded_save_params() + ["actor", "critic", "actor_target", "critic_target"] + return super()._excluded_save_params() + ["actor", "critic", "actor_target", "critic_target"] # noqa: RUF005 def _get_torch_save_params(self) -> Tuple[List[str], List[str]]: state_dicts = ["policy", "actor.optimizer", "critic.optimizer"] diff --git a/stable_baselines3/version.txt b/stable_baselines3/version.txt index 8daa30f..13ef2a8 100644 --- a/stable_baselines3/version.txt +++ b/stable_baselines3/version.txt @@ -1 +1 @@ -1.8.0a8 +1.8.0a9 diff --git a/tests/test_logger.py b/tests/test_logger.py index 94f1e21..a1d1052 100644 --- a/tests/test_logger.py +++ b/tests/test_logger.py @@ -234,7 +234,7 @@ def test_report_video_to_tensorboard(tmp_path, read_log, capsys): def is_moviepy_installed(): try: - import moviepy # noqa: F401 + import moviepy except ModuleNotFoundError: return False return True diff --git a/tests/test_utils.py b/tests/test_utils.py index dff80a0..0bab41e 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -441,7 +441,7 @@ def test_is_vectorized_observation(): # pass # All vectorized box_space = spaces.Box(-1, 1, shape=(2,)) - box_obs = np.ones((1,) + box_space.shape) + box_obs = np.ones((1, *box_space.shape)) assert is_vectorized_observation(box_obs, box_space) discrete_space = spaces.Discrete(2) @@ -485,13 +485,13 @@ def test_is_vectorized_observation(): # Vectorized with the wrong shape with pytest.raises(ValueError): discrete_obs = np.ones((1,), dtype=np.int8) - box_obs = np.ones((1, 2) + box_space.shape) + box_obs = np.ones((1, 2, *box_space.shape)) dict_obs = {"box": box_obs, "discrete": discrete_obs} is_vectorized_observation(dict_obs, dict_space) # Weird shape: error with pytest.raises(ValueError): - discrete_obs = np.ones((1,) + box_space.shape, dtype=np.int8) + discrete_obs = np.ones((1, *box_space.shape), dtype=np.int8) is_vectorized_observation(discrete_obs, discrete_space) # wrong shape @@ -506,7 +506,7 @@ def test_is_vectorized_observation(): # Almost good shape: one dimension too much for Discrete obs with pytest.raises(ValueError): - box_obs = np.ones((1,) + box_space.shape) + box_obs = np.ones((1, *box_space.shape)) discrete_obs = np.ones((1, 1), dtype=np.int8) dict_obs = {"box": box_obs, "discrete": discrete_obs} is_vectorized_observation(dict_obs, dict_space) diff --git a/tests/test_vec_envs.py b/tests/test_vec_envs.py index e661e1a..ae05947 100644 --- a/tests/test_vec_envs.py +++ b/tests/test_vec_envs.py @@ -361,10 +361,10 @@ def test_framestack_vecenv(): """Test that framestack environment stacks on desired axis""" image_space_shape = [12, 8, 3] - zero_acts = np.zeros([N_ENVS] + image_space_shape) + zero_acts = np.zeros([N_ENVS, *image_space_shape]) transposed_image_space_shape = image_space_shape[::-1] - transposed_zero_acts = np.zeros([N_ENVS] + transposed_image_space_shape) + transposed_zero_acts = np.zeros([N_ENVS, *transposed_image_space_shape]) def make_image_env(): return CustomGymEnv(