mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-09-15 22:10:25 +00:00
Fix Atari Roms download, enable RUF linting (#1379)
* Add extra no Atari and fix CI for forks * Enable ruff rules * Change to no roms
This commit is contained in:
parent
10e83865ec
commit
470771b5c2
21 changed files with 69 additions and 59 deletions
5
.github/workflows/ci.yml
vendored
5
.github/workflows/ci.yml
vendored
|
|
@ -14,7 +14,6 @@ jobs:
|
||||||
env:
|
env:
|
||||||
TERM: xterm-256color
|
TERM: xterm-256color
|
||||||
FORCE_COLOR: 1
|
FORCE_COLOR: 1
|
||||||
ATARI_ROMS: ${{ secrets.ATARI_ROMS }}
|
|
||||||
|
|
||||||
# Skip CI if [ci skip] in the commit message
|
# Skip CI if [ci skip] in the commit message
|
||||||
if: "! contains(toJSON(github.event.commits.*.message), '[ci skip]')"
|
if: "! contains(toJSON(github.event.commits.*.message), '[ci skip]')"
|
||||||
|
|
@ -37,11 +36,11 @@ jobs:
|
||||||
|
|
||||||
# Install Atari Roms
|
# Install Atari Roms
|
||||||
pip install autorom
|
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
|
base64 Roms.tar.gz.b64 --decode &> Roms.tar.gz
|
||||||
AutoROM --accept-license --source-file 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
|
# Use headless version
|
||||||
pip install opencv-python-headless
|
pip install opencv-python-headless
|
||||||
- name: Lint with ruff
|
- name: Lint with ruff
|
||||||
|
|
|
||||||
|
|
@ -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
|
- Moved from ``setup.cg`` to ``pyproject.toml`` configuration file
|
||||||
- Switched from ``flake8`` to ``ruff``
|
- Switched from ``flake8`` to ``ruff``
|
||||||
- Upgraded AutoROM to latest version
|
- Upgraded AutoROM to latest version
|
||||||
|
- Added ``extra_no_roms`` option for package installation without Atari Roms
|
||||||
|
|
||||||
Documentation:
|
Documentation:
|
||||||
^^^^^^^^^^^^^^
|
^^^^^^^^^^^^^^
|
||||||
|
|
|
||||||
|
|
@ -3,8 +3,8 @@
|
||||||
line-length = 127
|
line-length = 127
|
||||||
# Assume Python 3.7
|
# Assume Python 3.7
|
||||||
target-version = "py37"
|
target-version = "py37"
|
||||||
# TODO(antonin): activate "RUF" https://beta.ruff.rs/docs/rules/#ruff-specific-rules-ruf
|
# See https://beta.ruff.rs/docs/rules/
|
||||||
select = ["E", "F", "B", "UP", "C90"]
|
select = ["E", "F", "B", "UP", "C90", "RUF"]
|
||||||
ignore = []
|
ignore = []
|
||||||
|
|
||||||
[tool.ruff.per-file-ignores]
|
[tool.ruff.per-file-ignores]
|
||||||
|
|
|
||||||
40
setup.py
40
setup.py
|
|
@ -70,6 +70,29 @@ model = PPO("MlpPolicy", "CartPole-v1").learn(10_000)
|
||||||
|
|
||||||
""" # noqa:E501
|
""" # 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(
|
setup(
|
||||||
name="stable_baselines3",
|
name="stable_baselines3",
|
||||||
|
|
@ -119,21 +142,8 @@ setup(
|
||||||
# Copy button for code snippets
|
# Copy button for code snippets
|
||||||
"sphinx_copybutton",
|
"sphinx_copybutton",
|
||||||
],
|
],
|
||||||
"extra": [
|
"extra": extra_packages,
|
||||||
# For render
|
"extra_no_roms": extra_no_roms,
|
||||||
"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",
|
|
||||||
],
|
|
||||||
},
|
},
|
||||||
description="Pytorch version of Stable Baselines, implementations of reinforcement learning algorithms.",
|
description="Pytorch version of Stable Baselines, implementations of reinforcement learning algorithms.",
|
||||||
author="Antonin Raffin",
|
author="Antonin Raffin",
|
||||||
|
|
|
||||||
|
|
@ -156,7 +156,7 @@ class MaxAndSkipEnv(gym.Wrapper):
|
||||||
def __init__(self, env: gym.Env, skip: int = 4) -> None:
|
def __init__(self, env: gym.Env, skip: int = 4) -> None:
|
||||||
super().__init__(env)
|
super().__init__(env)
|
||||||
# most recent raw observations (for max pooling across time steps)
|
# 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
|
self._skip = skip
|
||||||
|
|
||||||
def step(self, action: int) -> GymStepReturn:
|
def step(self, action: int) -> GymStepReturn:
|
||||||
|
|
|
||||||
|
|
@ -67,7 +67,7 @@ class BaseBuffer(ABC):
|
||||||
"""
|
"""
|
||||||
shape = arr.shape
|
shape = arr.shape
|
||||||
if len(shape) < 3:
|
if len(shape) < 3:
|
||||||
shape = shape + (1,)
|
shape = (*shape, 1)
|
||||||
return arr.swapaxes(0, 1).reshape(shape[0] * shape[1], *shape[2:])
|
return arr.swapaxes(0, 1).reshape(shape[0] * shape[1], *shape[2:])
|
||||||
|
|
||||||
def size(self) -> int:
|
def size(self) -> int:
|
||||||
|
|
@ -199,13 +199,13 @@ class ReplayBuffer(BaseBuffer):
|
||||||
)
|
)
|
||||||
self.optimize_memory_usage = optimize_memory_usage
|
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:
|
if optimize_memory_usage:
|
||||||
# `observations` contains also the next observation
|
# `observations` contains also the next observation
|
||||||
self.next_observations = None
|
self.next_observations = None
|
||||||
else:
|
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)
|
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
|
# Reshape needed when using multiple envs with discrete observations
|
||||||
# as numpy cannot broadcast (n_discrete,) to (n_discrete, 1)
|
# as numpy cannot broadcast (n_discrete,) to (n_discrete, 1)
|
||||||
if isinstance(self.observation_space, spaces.Discrete):
|
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))
|
||||||
next_obs = next_obs.reshape((self.n_envs,) + self.obs_shape)
|
next_obs = next_obs.reshape((self.n_envs, *self.obs_shape))
|
||||||
|
|
||||||
# Same, for actions
|
# Same, for actions
|
||||||
action = action.reshape((self.n_envs, self.action_dim))
|
action = action.reshape((self.n_envs, self.action_dim))
|
||||||
|
|
@ -354,7 +354,7 @@ class RolloutBuffer(BaseBuffer):
|
||||||
self.reset()
|
self.reset()
|
||||||
|
|
||||||
def reset(self) -> None:
|
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.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.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)
|
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
|
# Reshape needed when using multiple envs with discrete observations
|
||||||
# as numpy cannot broadcast (n_discrete,) to (n_discrete, 1)
|
# as numpy cannot broadcast (n_discrete,) to (n_discrete, 1)
|
||||||
if isinstance(self.observation_space, spaces.Discrete):
|
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
|
# Same reshape, for actions
|
||||||
action = action.reshape((self.n_envs, self.action_dim))
|
action = action.reshape((self.n_envs, self.action_dim))
|
||||||
|
|
@ -528,11 +528,11 @@ class DictReplayBuffer(ReplayBuffer):
|
||||||
self.optimize_memory_usage = optimize_memory_usage
|
self.optimize_memory_usage = optimize_memory_usage
|
||||||
|
|
||||||
self.observations = {
|
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()
|
for key, _obs_shape in self.obs_shape.items()
|
||||||
}
|
}
|
||||||
self.next_observations = {
|
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()
|
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"
|
assert isinstance(self.obs_shape, dict), "DictRolloutBuffer must be used with Dict obs space only"
|
||||||
self.observations = {}
|
self.observations = {}
|
||||||
for key, obs_input_shape in self.obs_shape.items():
|
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.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.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)
|
self.returns = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32)
|
||||||
|
|
|
||||||
|
|
@ -554,7 +554,7 @@ class StopTrainingOnRewardThreshold(BaseCallback):
|
||||||
|
|
||||||
class EveryNTimesteps(EventCallback):
|
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 n_steps: Number of timesteps between two trigger.
|
||||||
:param callback: Callback that will be called
|
:param callback: Callback that will be called
|
||||||
|
|
|
||||||
|
|
@ -194,7 +194,7 @@ class ResultsWriter:
|
||||||
mode = "w" if override_existing else "a"
|
mode = "w" if override_existing else "a"
|
||||||
# Prevent newline issue on Windows, see GH issue #692
|
# Prevent newline issue on Windows, see GH issue #692
|
||||||
self.file_handler = open(filename, f"{mode}t", newline="\n")
|
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:
|
if override_existing:
|
||||||
self.file_handler.write(f"#{json.dumps(header)}\n")
|
self.file_handler.write(f"#{json.dumps(header)}\n")
|
||||||
self.logger.writeheader()
|
self.logger.writeheader()
|
||||||
|
|
|
||||||
|
|
@ -228,7 +228,7 @@ class BaseModel(nn.Module):
|
||||||
obs_ = np.array(obs)
|
obs_ = np.array(obs)
|
||||||
vectorized_env = vectorized_env or is_vectorized_observation(obs_, obs_space)
|
vectorized_env = vectorized_env or is_vectorized_observation(obs_, obs_space)
|
||||||
# Add batch dimension if needed
|
# 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):
|
elif is_image_space(self.observation_space):
|
||||||
# Handle the different cases for images
|
# Handle the different cases for images
|
||||||
|
|
@ -242,7 +242,7 @@ class BaseModel(nn.Module):
|
||||||
# Dict obs need to be handled separately
|
# Dict obs need to be handled separately
|
||||||
vectorized_env = is_vectorized_observation(observation, self.observation_space)
|
vectorized_env = is_vectorized_observation(observation, self.observation_space)
|
||||||
# Add batch dimension if needed
|
# 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)
|
observation = obs_as_tensor(observation, self.device)
|
||||||
return observation, vectorized_env
|
return observation, vectorized_env
|
||||||
|
|
@ -330,7 +330,7 @@ class BasePolicy(BaseModel, ABC):
|
||||||
with th.no_grad():
|
with th.no_grad():
|
||||||
actions = self._predict(observation, deterministic=deterministic)
|
actions = self._predict(observation, deterministic=deterministic)
|
||||||
# Convert to numpy, and reshape to the original action shape
|
# 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 isinstance(self.action_space, spaces.Box):
|
||||||
if self.squash_output:
|
if self.squash_output:
|
||||||
|
|
@ -608,7 +608,7 @@ class ActorCriticPolicy(BasePolicy):
|
||||||
distribution = self._get_action_dist_from_latent(latent_pi)
|
distribution = self._get_action_dist_from_latent(latent_pi)
|
||||||
actions = distribution.get_actions(deterministic=deterministic)
|
actions = distribution.get_actions(deterministic=deterministic)
|
||||||
log_prob = distribution.log_prob(actions)
|
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
|
return actions, values, log_prob
|
||||||
|
|
||||||
def extract_features(self, obs: th.Tensor) -> Union[th.Tensor, Tuple[th.Tensor, th.Tensor]]:
|
def extract_features(self, obs: th.Tensor) -> Union[th.Tensor, Tuple[th.Tensor, th.Tensor]]:
|
||||||
|
|
|
||||||
|
|
@ -25,7 +25,7 @@ def rolling_window(array: np.ndarray, window: int) -> np.ndarray:
|
||||||
:return: rolling window on the input array
|
:return: rolling window on the input array
|
||||||
"""
|
"""
|
||||||
shape = array.shape[:-1] + (array.shape[-1] - window + 1, window)
|
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)
|
return np.lib.stride_tricks.as_strided(array, shape=shape, strides=strides)
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -37,7 +37,7 @@ def recursive_getattr(obj: Any, attr: str, *args) -> Any:
|
||||||
def _getattr(obj: Any, attr: str) -> Any:
|
def _getattr(obj: Any, attr: str) -> Any:
|
||||||
return getattr(obj, attr, *args)
|
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:
|
def recursive_setattr(obj: Any, attr: str, val: Any) -> None:
|
||||||
|
|
|
||||||
|
|
@ -39,7 +39,7 @@ class DummyVecEnv(VecEnv):
|
||||||
obs_space = env.observation_space
|
obs_space = env.observation_space
|
||||||
self.keys, shapes, dtypes = obs_space_info(obs_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_dones = np.zeros((self.num_envs,), dtype=bool)
|
||||||
self.buf_rews = np.zeros((self.num_envs,), dtype=np.float32)
|
self.buf_rews = np.zeros((self.num_envs,), dtype=np.float32)
|
||||||
self.buf_infos = [{} for _ in range(self.num_envs)]
|
self.buf_infos = [{} for _ in range(self.num_envs)]
|
||||||
|
|
|
||||||
|
|
@ -56,7 +56,7 @@ class StackedObservations(Generic[TObs]):
|
||||||
low = np.repeat(observation_space.low, n_stack, axis=self.repeat_axis)
|
low = np.repeat(observation_space.low, n_stack, axis=self.repeat_axis)
|
||||||
high = np.repeat(observation_space.high, 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_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:
|
else:
|
||||||
raise TypeError(
|
raise TypeError(
|
||||||
f"StackedObservations only supports Box and Dict as observation spaces. {observation_space} was provided."
|
f"StackedObservations only supports Box and Dict as observation spaces. {observation_space} was provided."
|
||||||
|
|
|
||||||
|
|
@ -270,7 +270,7 @@ class DQN(OffPolicyAlgorithm):
|
||||||
)
|
)
|
||||||
|
|
||||||
def _excluded_save_params(self) -> List[str]:
|
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]]:
|
def _get_torch_save_params(self) -> Tuple[List[str], List[str]]:
|
||||||
state_dicts = ["policy", "policy.optimizer"]
|
state_dicts = ["policy", "policy.optimizer"]
|
||||||
|
|
|
||||||
|
|
@ -130,14 +130,14 @@ class HerReplayBuffer(DictReplayBuffer):
|
||||||
|
|
||||||
# input dimensions for buffer initialization
|
# input dimensions for buffer initialization
|
||||||
input_shape = {
|
input_shape = {
|
||||||
"observation": (self.env.num_envs,) + self.obs_shape,
|
"observation": (self.env.num_envs, *self.obs_shape),
|
||||||
"achieved_goal": (self.env.num_envs,) + self.goal_shape,
|
"achieved_goal": (self.env.num_envs, *self.goal_shape),
|
||||||
"desired_goal": (self.env.num_envs,) + self.goal_shape,
|
"desired_goal": (self.env.num_envs, *self.goal_shape),
|
||||||
"action": (self.action_dim,),
|
"action": (self.action_dim,),
|
||||||
"reward": (1,),
|
"reward": (1,),
|
||||||
"next_obs": (self.env.num_envs,) + self.obs_shape,
|
"next_obs": (self.env.num_envs, *self.obs_shape),
|
||||||
"next_achieved_goal": (self.env.num_envs,) + self.goal_shape,
|
"next_achieved_goal": (self.env.num_envs, *self.goal_shape),
|
||||||
"next_desired_goal": (self.env.num_envs,) + self.goal_shape,
|
"next_desired_goal": (self.env.num_envs, *self.goal_shape),
|
||||||
"done": (1,),
|
"done": (1,),
|
||||||
}
|
}
|
||||||
self._observation_keys = ["observation", "achieved_goal", "desired_goal"]
|
self._observation_keys = ["observation", "achieved_goal", "desired_goal"]
|
||||||
|
|
|
||||||
|
|
@ -304,7 +304,7 @@ class SAC(OffPolicyAlgorithm):
|
||||||
)
|
)
|
||||||
|
|
||||||
def _excluded_save_params(self) -> List[str]:
|
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]]:
|
def _get_torch_save_params(self) -> Tuple[List[str], List[str]]:
|
||||||
state_dicts = ["policy", "actor.optimizer", "critic.optimizer"]
|
state_dicts = ["policy", "actor.optimizer", "critic.optimizer"]
|
||||||
|
|
|
||||||
|
|
@ -218,7 +218,7 @@ class TD3(OffPolicyAlgorithm):
|
||||||
)
|
)
|
||||||
|
|
||||||
def _excluded_save_params(self) -> List[str]:
|
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]]:
|
def _get_torch_save_params(self) -> Tuple[List[str], List[str]]:
|
||||||
state_dicts = ["policy", "actor.optimizer", "critic.optimizer"]
|
state_dicts = ["policy", "actor.optimizer", "critic.optimizer"]
|
||||||
|
|
|
||||||
|
|
@ -1 +1 @@
|
||||||
1.8.0a8
|
1.8.0a9
|
||||||
|
|
|
||||||
|
|
@ -234,7 +234,7 @@ def test_report_video_to_tensorboard(tmp_path, read_log, capsys):
|
||||||
|
|
||||||
def is_moviepy_installed():
|
def is_moviepy_installed():
|
||||||
try:
|
try:
|
||||||
import moviepy # noqa: F401
|
import moviepy
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
return False
|
return False
|
||||||
return True
|
return True
|
||||||
|
|
|
||||||
|
|
@ -441,7 +441,7 @@ def test_is_vectorized_observation():
|
||||||
# pass
|
# pass
|
||||||
# All vectorized
|
# All vectorized
|
||||||
box_space = spaces.Box(-1, 1, shape=(2,))
|
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)
|
assert is_vectorized_observation(box_obs, box_space)
|
||||||
|
|
||||||
discrete_space = spaces.Discrete(2)
|
discrete_space = spaces.Discrete(2)
|
||||||
|
|
@ -485,13 +485,13 @@ def test_is_vectorized_observation():
|
||||||
# Vectorized with the wrong shape
|
# Vectorized with the wrong shape
|
||||||
with pytest.raises(ValueError):
|
with pytest.raises(ValueError):
|
||||||
discrete_obs = np.ones((1,), dtype=np.int8)
|
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}
|
dict_obs = {"box": box_obs, "discrete": discrete_obs}
|
||||||
is_vectorized_observation(dict_obs, dict_space)
|
is_vectorized_observation(dict_obs, dict_space)
|
||||||
|
|
||||||
# Weird shape: error
|
# Weird shape: error
|
||||||
with pytest.raises(ValueError):
|
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)
|
is_vectorized_observation(discrete_obs, discrete_space)
|
||||||
|
|
||||||
# wrong shape
|
# wrong shape
|
||||||
|
|
@ -506,7 +506,7 @@ def test_is_vectorized_observation():
|
||||||
|
|
||||||
# Almost good shape: one dimension too much for Discrete obs
|
# Almost good shape: one dimension too much for Discrete obs
|
||||||
with pytest.raises(ValueError):
|
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)
|
discrete_obs = np.ones((1, 1), dtype=np.int8)
|
||||||
dict_obs = {"box": box_obs, "discrete": discrete_obs}
|
dict_obs = {"box": box_obs, "discrete": discrete_obs}
|
||||||
is_vectorized_observation(dict_obs, dict_space)
|
is_vectorized_observation(dict_obs, dict_space)
|
||||||
|
|
|
||||||
|
|
@ -361,10 +361,10 @@ def test_framestack_vecenv():
|
||||||
"""Test that framestack environment stacks on desired axis"""
|
"""Test that framestack environment stacks on desired axis"""
|
||||||
|
|
||||||
image_space_shape = [12, 8, 3]
|
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_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():
|
def make_image_env():
|
||||||
return CustomGymEnv(
|
return CustomGymEnv(
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue