Upgrade code to Python 3.7+ syntax using pyupgrade (#887)

* Upgrade code to Python 3.7+ syntax

* Update changelog
This commit is contained in:
Antonin RAFFIN 2022-04-25 12:01:38 +02:00 committed by GitHub
parent 061841a314
commit a6f5049a99
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
44 changed files with 129 additions and 129 deletions

View file

@ -1,4 +1,3 @@
# -*- coding: utf-8 -*-
# #
# Configuration file for the Sphinx documentation builder. # Configuration file for the Sphinx documentation builder.
# #
@ -46,7 +45,7 @@ sys.modules.update((mod_name, Mock()) for mod_name in MOCK_MODULES)
# Read version from file # Read version from file
version_file = os.path.join(os.path.dirname(__file__), "../stable_baselines3", "version.txt") version_file = os.path.join(os.path.dirname(__file__), "../stable_baselines3", "version.txt")
with open(version_file, "r") as file_handler: with open(version_file) as file_handler:
__version__ = file_handler.read().strip() __version__ = file_handler.read().strip()
# -- Project information ----------------------------------------------------- # -- Project information -----------------------------------------------------

View file

@ -4,7 +4,7 @@ Changelog
========== ==========
Release 1.5.1a4 (WIP) Release 1.5.1a5 (WIP)
--------------------------- ---------------------------
Breaking Changes: Breaking Changes:
@ -31,6 +31,7 @@ Deprecations:
Others: Others:
^^^^^^^ ^^^^^^^
- Upgraded to Python 3.7+ syntax using ``pyupgrade``
Documentation: Documentation:
^^^^^^^^^^^^^^ ^^^^^^^^^^^^^^

View file

@ -2,7 +2,7 @@ import os
from setuptools import find_packages, setup from setuptools import find_packages, setup
with open(os.path.join("stable_baselines3", "version.txt"), "r") as file_handler: with open(os.path.join("stable_baselines3", "version.txt")) as file_handler:
__version__ = file_handler.read().strip() __version__ = file_handler.read().strip()

View file

@ -11,7 +11,7 @@ from stable_baselines3.td3 import TD3
# Read version from file # Read version from file
version_file = os.path.join(os.path.dirname(__file__), "version.txt") version_file = os.path.join(os.path.dirname(__file__), "version.txt")
with open(version_file, "r") as file_handler: with open(version_file) as file_handler:
__version__ = file_handler.read().strip() __version__ = file_handler.read().strip()

View file

@ -82,7 +82,7 @@ class A2C(OnPolicyAlgorithm):
_init_setup_model: bool = True, _init_setup_model: bool = True,
): ):
super(A2C, self).__init__( super().__init__(
policy, policy,
env, env,
learning_rate=learning_rate, learning_rate=learning_rate,
@ -194,7 +194,7 @@ class A2C(OnPolicyAlgorithm):
reset_num_timesteps: bool = True, reset_num_timesteps: bool = True,
) -> "A2C": ) -> "A2C":
return super(A2C, self).learn( return super().learn(
total_timesteps=total_timesteps, total_timesteps=total_timesteps,
callback=callback, callback=callback,
log_interval=log_interval, log_interval=log_interval,

View file

@ -245,4 +245,4 @@ class AtariWrapper(gym.Wrapper):
if clip_reward: if clip_reward:
env = ClipRewardEnv(env) env = ClipRewardEnv(env)
super(AtariWrapper, self).__init__(env) super().__init__(env)

View file

@ -42,7 +42,7 @@ class BaseBuffer(ABC):
device: Union[th.device, str] = "cpu", device: Union[th.device, str] = "cpu",
n_envs: int = 1, n_envs: int = 1,
): ):
super(BaseBuffer, self).__init__() super().__init__()
self.buffer_size = buffer_size self.buffer_size = buffer_size
self.observation_space = observation_space self.observation_space = observation_space
self.action_space = action_space self.action_space = action_space
@ -179,7 +179,7 @@ class ReplayBuffer(BaseBuffer):
optimize_memory_usage: bool = False, optimize_memory_usage: bool = False,
handle_timeout_termination: bool = True, handle_timeout_termination: bool = True,
): ):
super(ReplayBuffer, self).__init__(buffer_size, observation_space, action_space, device, n_envs=n_envs) super().__init__(buffer_size, observation_space, action_space, device, n_envs=n_envs)
# Adjust buffer size # Adjust buffer size
self.buffer_size = max(buffer_size // n_envs, 1) self.buffer_size = max(buffer_size // n_envs, 1)
@ -339,7 +339,7 @@ class RolloutBuffer(BaseBuffer):
n_envs: int = 1, n_envs: int = 1,
): ):
super(RolloutBuffer, self).__init__(buffer_size, observation_space, action_space, device, n_envs=n_envs) super().__init__(buffer_size, observation_space, action_space, device, n_envs=n_envs)
self.gae_lambda = gae_lambda self.gae_lambda = gae_lambda
self.gamma = gamma self.gamma = gamma
self.observations, self.actions, self.rewards, self.advantages = None, None, None, None self.observations, self.actions, self.rewards, self.advantages = None, None, None, None
@ -358,7 +358,7 @@ class RolloutBuffer(BaseBuffer):
self.log_probs = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32) self.log_probs = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32)
self.advantages = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32) self.advantages = np.zeros((self.buffer_size, self.n_envs), dtype=np.float32)
self.generator_ready = False self.generator_ready = False
super(RolloutBuffer, self).reset() super().reset()
def compute_returns_and_advantage(self, last_values: th.Tensor, dones: np.ndarray) -> None: def compute_returns_and_advantage(self, last_values: th.Tensor, dones: np.ndarray) -> None:
""" """

View file

@ -19,7 +19,7 @@ class BaseCallback(ABC):
""" """
def __init__(self, verbose: int = 0): def __init__(self, verbose: int = 0):
super(BaseCallback, self).__init__() super().__init__()
# The RL model # The RL model
self.model = None # type: Optional[base_class.BaseAlgorithm] self.model = None # type: Optional[base_class.BaseAlgorithm]
# An alias for self.model.get_env(), the environment used for training # An alias for self.model.get_env(), the environment used for training
@ -127,14 +127,14 @@ class EventCallback(BaseCallback):
""" """
def __init__(self, callback: Optional[BaseCallback] = None, verbose: int = 0): def __init__(self, callback: Optional[BaseCallback] = None, verbose: int = 0):
super(EventCallback, self).__init__(verbose=verbose) super().__init__(verbose=verbose)
self.callback = callback self.callback = callback
# Give access to the parent # Give access to the parent
if callback is not None: if callback is not None:
self.callback.parent = self self.callback.parent = self
def init_callback(self, model: "base_class.BaseAlgorithm") -> None: def init_callback(self, model: "base_class.BaseAlgorithm") -> None:
super(EventCallback, self).init_callback(model) super().init_callback(model)
if self.callback is not None: if self.callback is not None:
self.callback.init_callback(self.model) self.callback.init_callback(self.model)
@ -169,7 +169,7 @@ class CallbackList(BaseCallback):
""" """
def __init__(self, callbacks: List[BaseCallback]): def __init__(self, callbacks: List[BaseCallback]):
super(CallbackList, self).__init__() super().__init__()
assert isinstance(callbacks, list) assert isinstance(callbacks, list)
self.callbacks = callbacks self.callbacks = callbacks
@ -228,7 +228,7 @@ class CheckpointCallback(BaseCallback):
""" """
def __init__(self, save_freq: int, save_path: str, name_prefix: str = "rl_model", verbose: int = 0): def __init__(self, save_freq: int, save_path: str, name_prefix: str = "rl_model", verbose: int = 0):
super(CheckpointCallback, self).__init__(verbose) super().__init__(verbose)
self.save_freq = save_freq self.save_freq = save_freq
self.save_path = save_path self.save_path = save_path
self.name_prefix = name_prefix self.name_prefix = name_prefix
@ -256,7 +256,7 @@ class ConvertCallback(BaseCallback):
""" """
def __init__(self, callback: Callable[[Dict[str, Any], Dict[str, Any]], bool], verbose: int = 0): def __init__(self, callback: Callable[[Dict[str, Any], Dict[str, Any]], bool], verbose: int = 0):
super(ConvertCallback, self).__init__(verbose) super().__init__(verbose)
self.callback = callback self.callback = callback
def _on_step(self) -> bool: def _on_step(self) -> bool:
@ -307,7 +307,7 @@ class EvalCallback(EventCallback):
verbose: int = 1, verbose: int = 1,
warn: bool = True, warn: bool = True,
): ):
super(EvalCallback, self).__init__(callback_after_eval, verbose=verbose) super().__init__(callback_after_eval, verbose=verbose)
self.callback_on_new_best = callback_on_new_best self.callback_on_new_best = callback_on_new_best
if self.callback_on_new_best is not None: if self.callback_on_new_best is not None:
@ -480,7 +480,7 @@ class StopTrainingOnRewardThreshold(BaseCallback):
""" """
def __init__(self, reward_threshold: float, verbose: int = 0): def __init__(self, reward_threshold: float, verbose: int = 0):
super(StopTrainingOnRewardThreshold, self).__init__(verbose=verbose) super().__init__(verbose=verbose)
self.reward_threshold = reward_threshold self.reward_threshold = reward_threshold
def _on_step(self) -> bool: def _on_step(self) -> bool:
@ -505,7 +505,7 @@ class EveryNTimesteps(EventCallback):
""" """
def __init__(self, n_steps: int, callback: BaseCallback): def __init__(self, n_steps: int, callback: BaseCallback):
super(EveryNTimesteps, self).__init__(callback) super().__init__(callback)
self.n_steps = n_steps self.n_steps = n_steps
self.last_time_trigger = 0 self.last_time_trigger = 0
@ -528,7 +528,7 @@ class StopTrainingOnMaxEpisodes(BaseCallback):
""" """
def __init__(self, max_episodes: int, verbose: int = 0): def __init__(self, max_episodes: int, verbose: int = 0):
super(StopTrainingOnMaxEpisodes, self).__init__(verbose=verbose) super().__init__(verbose=verbose)
self.max_episodes = max_episodes self.max_episodes = max_episodes
self._total_max_episodes = max_episodes self._total_max_episodes = max_episodes
self.n_episodes = 0 self.n_episodes = 0
@ -573,7 +573,7 @@ class StopTrainingOnNoModelImprovement(BaseCallback):
""" """
def __init__(self, max_no_improvement_evals: int, min_evals: int = 0, verbose: int = 0): def __init__(self, max_no_improvement_evals: int, min_evals: int = 0, verbose: int = 0):
super(StopTrainingOnNoModelImprovement, self).__init__(verbose=verbose) super().__init__(verbose=verbose)
self.max_no_improvement_evals = max_no_improvement_evals self.max_no_improvement_evals = max_no_improvement_evals
self.min_evals = min_evals self.min_evals = min_evals
self.last_best_mean_reward = -np.inf self.last_best_mean_reward = -np.inf

View file

@ -16,7 +16,7 @@ class Distribution(ABC):
"""Abstract base class for distributions.""" """Abstract base class for distributions."""
def __init__(self): def __init__(self):
super(Distribution, self).__init__() super().__init__()
self.distribution = None self.distribution = None
@abstractmethod @abstractmethod
@ -120,7 +120,7 @@ class DiagGaussianDistribution(Distribution):
""" """
def __init__(self, action_dim: int): def __init__(self, action_dim: int):
super(DiagGaussianDistribution, self).__init__() super().__init__()
self.action_dim = action_dim self.action_dim = action_dim
self.mean_actions = None self.mean_actions = None
self.log_std = None self.log_std = None
@ -201,13 +201,13 @@ class SquashedDiagGaussianDistribution(DiagGaussianDistribution):
""" """
def __init__(self, action_dim: int, epsilon: float = 1e-6): def __init__(self, action_dim: int, epsilon: float = 1e-6):
super(SquashedDiagGaussianDistribution, self).__init__(action_dim) super().__init__(action_dim)
# Avoid NaN (prevents division by zero or log of zero) # Avoid NaN (prevents division by zero or log of zero)
self.epsilon = epsilon self.epsilon = epsilon
self.gaussian_actions = None self.gaussian_actions = None
def proba_distribution(self, mean_actions: th.Tensor, log_std: th.Tensor) -> "SquashedDiagGaussianDistribution": def proba_distribution(self, mean_actions: th.Tensor, log_std: th.Tensor) -> "SquashedDiagGaussianDistribution":
super(SquashedDiagGaussianDistribution, self).proba_distribution(mean_actions, log_std) super().proba_distribution(mean_actions, log_std)
return self return self
def log_prob(self, actions: th.Tensor, gaussian_actions: Optional[th.Tensor] = None) -> th.Tensor: def log_prob(self, actions: th.Tensor, gaussian_actions: Optional[th.Tensor] = None) -> th.Tensor:
@ -219,7 +219,7 @@ class SquashedDiagGaussianDistribution(DiagGaussianDistribution):
gaussian_actions = TanhBijector.inverse(actions) gaussian_actions = TanhBijector.inverse(actions)
# Log likelihood for a Gaussian distribution # Log likelihood for a Gaussian distribution
log_prob = super(SquashedDiagGaussianDistribution, self).log_prob(gaussian_actions) log_prob = super().log_prob(gaussian_actions)
# Squash correction (from original SAC implementation) # Squash correction (from original SAC implementation)
# this comes from the fact that tanh is bijective and differentiable # this comes from the fact that tanh is bijective and differentiable
log_prob -= th.sum(th.log(1 - actions**2 + self.epsilon), dim=1) log_prob -= th.sum(th.log(1 - actions**2 + self.epsilon), dim=1)
@ -254,7 +254,7 @@ class CategoricalDistribution(Distribution):
""" """
def __init__(self, action_dim: int): def __init__(self, action_dim: int):
super(CategoricalDistribution, self).__init__() super().__init__()
self.action_dim = action_dim self.action_dim = action_dim
def proba_distribution_net(self, latent_dim: int) -> nn.Module: def proba_distribution_net(self, latent_dim: int) -> nn.Module:
@ -305,7 +305,7 @@ class MultiCategoricalDistribution(Distribution):
""" """
def __init__(self, action_dims: List[int]): def __init__(self, action_dims: List[int]):
super(MultiCategoricalDistribution, self).__init__() super().__init__()
self.action_dims = action_dims self.action_dims = action_dims
def proba_distribution_net(self, latent_dim: int) -> nn.Module: def proba_distribution_net(self, latent_dim: int) -> nn.Module:
@ -360,7 +360,7 @@ class BernoulliDistribution(Distribution):
""" """
def __init__(self, action_dims: int): def __init__(self, action_dims: int):
super(BernoulliDistribution, self).__init__() super().__init__()
self.action_dims = action_dims self.action_dims = action_dims
def proba_distribution_net(self, latent_dim: int) -> nn.Module: def proba_distribution_net(self, latent_dim: int) -> nn.Module:
@ -433,7 +433,7 @@ class StateDependentNoiseDistribution(Distribution):
learn_features: bool = False, learn_features: bool = False,
epsilon: float = 1e-6, epsilon: float = 1e-6,
): ):
super(StateDependentNoiseDistribution, self).__init__() super().__init__()
self.action_dim = action_dim self.action_dim = action_dim
self.latent_sde_dim = None self.latent_sde_dim = None
self.mean_actions = None self.mean_actions = None
@ -597,7 +597,7 @@ class StateDependentNoiseDistribution(Distribution):
return actions, log_prob return actions, log_prob
class TanhBijector(object): class TanhBijector:
""" """
Bijective transformation of a probability distribution Bijective transformation of a probability distribution
using a squashing function (tanh) using a squashing function (tanh)
@ -607,7 +607,7 @@ class TanhBijector(object):
""" """
def __init__(self, epsilon: float = 1e-6): def __init__(self, epsilon: float = 1e-6):
super(TanhBijector, self).__init__() super().__init__()
self.epsilon = epsilon self.epsilon = epsilon
@staticmethod @staticmethod

View file

@ -36,7 +36,7 @@ class BitFlippingEnv(GoalEnv):
image_obs_space: bool = False, image_obs_space: bool = False,
channel_first: bool = True, channel_first: bool = True,
): ):
super(BitFlippingEnv, self).__init__() super().__init__()
# Shape of the observation when using image space # Shape of the observation when using image space
self.image_shape = (1, 36, 36) if channel_first else (36, 36, 1) self.image_shape = (1, 36, 36) if channel_first else (36, 36, 1)
# The achieved goal is determined by the current state # The achieved goal is determined by the current state
@ -115,7 +115,7 @@ class BitFlippingEnv(GoalEnv):
if self.discrete_obs_space: if self.discrete_obs_space:
# The internal state is the binary representation of the # The internal state is the binary representation of the
# observed one # observed one
return int(sum([state[i] * 2**i for i in range(len(state))])) return int(sum(state[i] * 2**i for i in range(len(state))))
if self.image_obs_space: if self.image_obs_space:
size = np.prod(self.image_shape) size = np.prod(self.image_shape)
@ -135,7 +135,7 @@ class BitFlippingEnv(GoalEnv):
if isinstance(state, int): if isinstance(state, int):
state = np.array(state).reshape(batch_size, -1) state = np.array(state).reshape(batch_size, -1)
# Convert to binary representation # Convert to binary representation
state = (((state[:, :] & (1 << np.arange(len(self.state))))) > 0).astype(int) state = ((state[:, :] & (1 << np.arange(len(self.state)))) > 0).astype(int)
elif self.image_obs_space: elif self.image_obs_space:
state = state.reshape(batch_size, -1)[:, : len(self.state)] / 255 state = state.reshape(batch_size, -1)[:, : len(self.state)] / 255
else: else:

View file

@ -42,7 +42,7 @@ class SimpleMultiObsEnv(gym.Env):
discrete_actions: bool = True, discrete_actions: bool = True,
channel_last: bool = True, channel_last: bool = True,
): ):
super(SimpleMultiObsEnv, self).__init__() super().__init__()
self.vector_size = 5 self.vector_size = 5
if channel_last: if channel_last:

View file

@ -24,7 +24,7 @@ ERROR = 40
DISABLED = 50 DISABLED = 50
class Video(object): class Video:
""" """
Video data class storing the video frames and the frame per seconds Video data class storing the video frames and the frame per seconds
@ -37,7 +37,7 @@ class Video(object):
self.fps = fps self.fps = fps
class Figure(object): class Figure:
""" """
Figure data class storing a matplotlib figure and whether to close the figure after logging it Figure data class storing a matplotlib figure and whether to close the figure after logging it
@ -50,7 +50,7 @@ class Figure(object):
self.close = close self.close = close
class Image(object): class Image:
""" """
Image data class storing an image and data format Image data class storing an image and data format
@ -80,13 +80,13 @@ class FormatUnsupportedError(NotImplementedError):
format_str = f"formats {', '.join(unsupported_formats)} are" format_str = f"formats {', '.join(unsupported_formats)} are"
else: else:
format_str = f"format {unsupported_formats[0]} is" format_str = f"format {unsupported_formats[0]} is"
super(FormatUnsupportedError, self).__init__( super().__init__(
f"The {format_str} not supported for the {value_description} value logged.\n" f"The {format_str} not supported for the {value_description} value logged.\n"
f"You can exclude formats via the `exclude` parameter of the logger's `record` function." f"You can exclude formats via the `exclude` parameter of the logger's `record` function."
) )
class KVWriter(object): class KVWriter:
""" """
Key Value writer Key Value writer
""" """
@ -108,7 +108,7 @@ class KVWriter(object):
raise NotImplementedError raise NotImplementedError
class SeqWriter(object): class SeqWriter:
""" """
sequence writer sequence writer
""" """
@ -427,7 +427,7 @@ def make_output_format(_format: str, log_dir: str, log_suffix: str = "") -> KVWr
# ================================================================ # ================================================================
class Logger(object): class Logger:
""" """
The logger class. The logger class.
@ -623,7 +623,7 @@ def read_json(filename: str) -> pandas.DataFrame:
:return: the data in the json :return: the data in the json
""" """
data = [] data = []
with open(filename, "rt") as file_handler: with open(filename) as file_handler:
for line in file_handler: for line in file_handler:
data.append(json.loads(line)) data.append(json.loads(line))
return pandas.DataFrame(data) return pandas.DataFrame(data)

View file

@ -36,7 +36,7 @@ class Monitor(gym.Wrapper):
reset_keywords: Tuple[str, ...] = (), reset_keywords: Tuple[str, ...] = (),
info_keywords: Tuple[str, ...] = (), info_keywords: Tuple[str, ...] = (),
): ):
super(Monitor, self).__init__(env=env) super().__init__(env=env)
self.t_start = time.time() self.t_start = time.time()
if filename is not None: if filename is not None:
self.results_writer = ResultsWriter( self.results_writer = ResultsWriter(
@ -110,7 +110,7 @@ class Monitor(gym.Wrapper):
""" """
Closes the environment Closes the environment
""" """
super(Monitor, self).close() super().close()
if self.results_writer is not None: if self.results_writer is not None:
self.results_writer.close() self.results_writer.close()
@ -224,7 +224,7 @@ def load_results(path: str) -> pandas.DataFrame:
raise LoadMonitorResultsError(f"No monitor files of the form *{Monitor.EXT} found in {path}") raise LoadMonitorResultsError(f"No monitor files of the form *{Monitor.EXT} found in {path}")
data_frames, headers = [], [] data_frames, headers = [], []
for file_name in monitor_files: for file_name in monitor_files:
with open(file_name, "rt") as file_handler: with open(file_name) as file_handler:
first_line = file_handler.readline() first_line = file_handler.readline()
assert first_line[0] == "#" assert first_line[0] == "#"
header = json.loads(first_line[1:]) header = json.loads(first_line[1:])

View file

@ -11,7 +11,7 @@ class ActionNoise(ABC):
""" """
def __init__(self): def __init__(self):
super(ActionNoise, self).__init__() super().__init__()
def reset(self) -> None: def reset(self) -> None:
""" """
@ -35,7 +35,7 @@ class NormalActionNoise(ActionNoise):
def __init__(self, mean: np.ndarray, sigma: np.ndarray): def __init__(self, mean: np.ndarray, sigma: np.ndarray):
self._mu = mean self._mu = mean
self._sigma = sigma self._sigma = sigma
super(NormalActionNoise, self).__init__() super().__init__()
def __call__(self) -> np.ndarray: def __call__(self) -> np.ndarray:
return np.random.normal(self._mu, self._sigma) return np.random.normal(self._mu, self._sigma)
@ -72,7 +72,7 @@ class OrnsteinUhlenbeckActionNoise(ActionNoise):
self.initial_noise = initial_noise self.initial_noise = initial_noise
self.noise_prev = np.zeros_like(self._mu) self.noise_prev = np.zeros_like(self._mu)
self.reset() self.reset()
super(OrnsteinUhlenbeckActionNoise, self).__init__() super().__init__()
def __call__(self) -> np.ndarray: def __call__(self) -> np.ndarray:
noise = ( noise = (

View file

@ -102,7 +102,7 @@ class OffPolicyAlgorithm(BaseAlgorithm):
supported_action_spaces: Optional[Tuple[gym.spaces.Space, ...]] = None, supported_action_spaces: Optional[Tuple[gym.spaces.Space, ...]] = None,
): ):
super(OffPolicyAlgorithm, self).__init__( super().__init__(
policy=policy, policy=policy,
env=env, env=env,
learning_rate=learning_rate, learning_rate=learning_rate,

View file

@ -72,7 +72,7 @@ class OnPolicyAlgorithm(BaseAlgorithm):
supported_action_spaces: Optional[Tuple[gym.spaces.Space, ...]] = None, supported_action_spaces: Optional[Tuple[gym.spaces.Space, ...]] = None,
): ):
super(OnPolicyAlgorithm, self).__init__( super().__init__(
policy=policy, policy=policy,
env=env, env=env,
learning_rate=learning_rate, learning_rate=learning_rate,

View file

@ -67,7 +67,7 @@ class BaseModel(nn.Module, ABC):
optimizer_class: Type[th.optim.Optimizer] = th.optim.Adam, optimizer_class: Type[th.optim.Optimizer] = th.optim.Adam,
optimizer_kwargs: Optional[Dict[str, Any]] = None, optimizer_kwargs: Optional[Dict[str, Any]] = None,
): ):
super(BaseModel, self).__init__() super().__init__()
if optimizer_kwargs is None: if optimizer_kwargs is None:
optimizer_kwargs = {} optimizer_kwargs = {}
@ -267,7 +267,7 @@ class BasePolicy(BaseModel):
""" """
def __init__(self, *args, squash_output: bool = False, **kwargs): def __init__(self, *args, squash_output: bool = False, **kwargs):
super(BasePolicy, self).__init__(*args, **kwargs) super().__init__(*args, **kwargs)
self._squash_output = squash_output self._squash_output = squash_output
@staticmethod @staticmethod
@ -437,7 +437,7 @@ class ActorCriticPolicy(BasePolicy):
if optimizer_class == th.optim.Adam: if optimizer_class == th.optim.Adam:
optimizer_kwargs["eps"] = 1e-5 optimizer_kwargs["eps"] = 1e-5
super(ActorCriticPolicy, self).__init__( super().__init__(
observation_space, observation_space,
action_space, action_space,
features_extractor_class, features_extractor_class,
@ -724,7 +724,7 @@ class ActorCriticCnnPolicy(ActorCriticPolicy):
optimizer_class: Type[th.optim.Optimizer] = th.optim.Adam, optimizer_class: Type[th.optim.Optimizer] = th.optim.Adam,
optimizer_kwargs: Optional[Dict[str, Any]] = None, optimizer_kwargs: Optional[Dict[str, Any]] = None,
): ):
super(ActorCriticCnnPolicy, self).__init__( super().__init__(
observation_space, observation_space,
action_space, action_space,
lr_schedule, lr_schedule,
@ -799,7 +799,7 @@ class MultiInputActorCriticPolicy(ActorCriticPolicy):
optimizer_class: Type[th.optim.Optimizer] = th.optim.Adam, optimizer_class: Type[th.optim.Optimizer] = th.optim.Adam,
optimizer_kwargs: Optional[Dict[str, Any]] = None, optimizer_kwargs: Optional[Dict[str, Any]] = None,
): ):
super(MultiInputActorCriticPolicy, self).__init__( super().__init__(
observation_space, observation_space,
action_space, action_space,
lr_schedule, lr_schedule,

View file

@ -3,7 +3,7 @@ from typing import Tuple, Union
import numpy as np import numpy as np
class RunningMeanStd(object): class RunningMeanStd:
def __init__(self, epsilon: float = 1e-4, shape: Tuple[int, ...] = ()): def __init__(self, epsilon: float = 1e-4, shape: Tuple[int, ...] = ()):
""" """
Calulates the running mean and std of a data stream Calulates the running mean and std of a data stream

View file

@ -54,21 +54,21 @@ class RMSpropTFLike(Optimizer):
centered: bool = False, centered: bool = False,
): ):
if not 0.0 <= lr: if not 0.0 <= lr:
raise ValueError("Invalid learning rate: {}".format(lr)) raise ValueError(f"Invalid learning rate: {lr}")
if not 0.0 <= eps: if not 0.0 <= eps:
raise ValueError("Invalid epsilon value: {}".format(eps)) raise ValueError(f"Invalid epsilon value: {eps}")
if not 0.0 <= momentum: if not 0.0 <= momentum:
raise ValueError("Invalid momentum value: {}".format(momentum)) raise ValueError(f"Invalid momentum value: {momentum}")
if not 0.0 <= weight_decay: if not 0.0 <= weight_decay:
raise ValueError("Invalid weight_decay value: {}".format(weight_decay)) raise ValueError(f"Invalid weight_decay value: {weight_decay}")
if not 0.0 <= alpha: if not 0.0 <= alpha:
raise ValueError("Invalid alpha value: {}".format(alpha)) raise ValueError(f"Invalid alpha value: {alpha}")
defaults = dict(lr=lr, momentum=momentum, alpha=alpha, eps=eps, centered=centered, weight_decay=weight_decay) defaults = dict(lr=lr, momentum=momentum, alpha=alpha, eps=eps, centered=centered, weight_decay=weight_decay)
super(RMSpropTFLike, self).__init__(params, defaults) super().__init__(params, defaults)
def __setstate__(self, state: Dict[str, Any]) -> None: def __setstate__(self, state: Dict[str, Any]) -> None:
super(RMSpropTFLike, self).__setstate__(state) super().__setstate__(state)
for group in self.param_groups: for group in self.param_groups:
group.setdefault("momentum", 0) group.setdefault("momentum", 0)
group.setdefault("centered", False) group.setdefault("centered", False)

View file

@ -19,7 +19,7 @@ class BaseFeaturesExtractor(nn.Module):
""" """
def __init__(self, observation_space: gym.Space, features_dim: int = 0): def __init__(self, observation_space: gym.Space, features_dim: int = 0):
super(BaseFeaturesExtractor, self).__init__() super().__init__()
assert features_dim > 0 assert features_dim > 0
self._observation_space = observation_space self._observation_space = observation_space
self._features_dim = features_dim self._features_dim = features_dim
@ -41,7 +41,7 @@ class FlattenExtractor(BaseFeaturesExtractor):
""" """
def __init__(self, observation_space: gym.Space): def __init__(self, observation_space: gym.Space):
super(FlattenExtractor, self).__init__(observation_space, get_flattened_obs_dim(observation_space)) super().__init__(observation_space, get_flattened_obs_dim(observation_space))
self.flatten = nn.Flatten() self.flatten = nn.Flatten()
def forward(self, observations: th.Tensor) -> th.Tensor: def forward(self, observations: th.Tensor) -> th.Tensor:
@ -61,7 +61,7 @@ class NatureCNN(BaseFeaturesExtractor):
""" """
def __init__(self, observation_space: gym.spaces.Box, features_dim: int = 512): def __init__(self, observation_space: gym.spaces.Box, features_dim: int = 512):
super(NatureCNN, self).__init__(observation_space, features_dim) super().__init__(observation_space, features_dim)
# We assume CxHxW images (channels first) # We assume CxHxW images (channels first)
# Re-ordering will be done by pre-preprocessing or wrapper # Re-ordering will be done by pre-preprocessing or wrapper
assert is_image_space(observation_space, check_channels=False), ( assert is_image_space(observation_space, check_channels=False), (
@ -169,7 +169,7 @@ class MlpExtractor(nn.Module):
activation_fn: Type[nn.Module], activation_fn: Type[nn.Module],
device: Union[th.device, str] = "auto", device: Union[th.device, str] = "auto",
): ):
super(MlpExtractor, self).__init__() super().__init__()
device = get_device(device) device = get_device(device)
shared_net, policy_net, value_net = [], [], [] shared_net, policy_net, value_net = [], [], []
policy_only_layers = [] # Layer sizes of the network that only belongs to the policy network policy_only_layers = [] # Layer sizes of the network that only belongs to the policy network
@ -250,7 +250,7 @@ class CombinedExtractor(BaseFeaturesExtractor):
def __init__(self, observation_space: gym.spaces.Dict, cnn_output_dim: int = 256): def __init__(self, observation_space: gym.spaces.Dict, cnn_output_dim: int = 256):
# TODO we do not know features-dim here before going over all the items, so put something there. This is dirty! # TODO we do not know features-dim here before going over all the items, so put something there. This is dirty!
super(CombinedExtractor, self).__init__(observation_space, features_dim=1) super().__init__(observation_space, features_dim=1)
extractors = {} extractors = {}

View file

@ -7,7 +7,7 @@ from gym import spaces
from stable_baselines3.common.preprocessing import is_image_space, is_image_space_channels_first from stable_baselines3.common.preprocessing import is_image_space, is_image_space_channels_first
class StackedObservations(object): class StackedObservations:
""" """
Frame stacking wrapper for data. Frame stacking wrapper for data.

View file

@ -217,6 +217,6 @@ def _flatten_obs(obs: Union[List[VecEnvObs], Tuple[VecEnvObs]], space: gym.space
elif isinstance(space, gym.spaces.Tuple): elif isinstance(space, gym.spaces.Tuple):
assert isinstance(obs[0], tuple), "non-tuple observation for environment with Tuple observation space" assert isinstance(obs[0], tuple), "non-tuple observation for environment with Tuple observation space"
obs_len = len(space.spaces) obs_len = len(space.spaces)
return tuple((np.stack([o[i] for o in obs]) for i in range(obs_len))) return tuple(np.stack([o[i] for o in obs]) for i in range(obs_len))
else: else:
return np.stack(obs) return np.stack(obs)

View file

@ -37,7 +37,7 @@ def dict_to_obs(obs_space: gym.spaces.Space, obs_dict: Dict[Any, np.ndarray]) ->
return obs_dict return obs_dict
elif isinstance(obs_space, gym.spaces.Tuple): elif isinstance(obs_space, gym.spaces.Tuple):
assert len(obs_dict) == len(obs_space.spaces), "size of observation does not match size of observation space" assert len(obs_dict) == len(obs_space.spaces), "size of observation does not match size of observation space"
return tuple((obs_dict[i] for i in range(len(obs_space.spaces)))) return tuple(obs_dict[i] for i in range(len(obs_space.spaces)))
else: else:
assert set(obs_dict.keys()) == {None}, "multiple observation keys for unstructured observation space" assert set(obs_dict.keys()) == {None}, "multiple observation keys for unstructured observation space"
return obs_dict[None] return obs_dict[None]

View file

@ -26,7 +26,7 @@ class VecTransposeImage(VecEnvWrapper):
self.skip = skip self.skip = skip
# Do nothing # Do nothing
if skip: if skip:
super(VecTransposeImage, self).__init__(venv) super().__init__(venv)
return return
if isinstance(venv.observation_space, spaces.dict.Dict): if isinstance(venv.observation_space, spaces.dict.Dict):
@ -39,7 +39,7 @@ class VecTransposeImage(VecEnvWrapper):
observation_space.spaces[key] = self.transpose_space(space, key) observation_space.spaces[key] = self.transpose_space(space, key)
else: else:
observation_space = self.transpose_space(venv.observation_space) observation_space = self.transpose_space(venv.observation_space)
super(VecTransposeImage, self).__init__(venv, observation_space=observation_space) super().__init__(venv, observation_space=observation_space)
@staticmethod @staticmethod
def transpose_space(observation_space: spaces.Box, key: str = "") -> spaces.Box: def transpose_space(observation_space: spaces.Box, key: str = "") -> spaces.Box:

View file

@ -78,7 +78,7 @@ class DDPG(TD3):
_init_setup_model: bool = True, _init_setup_model: bool = True,
): ):
super(DDPG, self).__init__( super().__init__(
policy=policy, policy=policy,
env=env, env=env,
learning_rate=learning_rate, learning_rate=learning_rate,
@ -127,7 +127,7 @@ class DDPG(TD3):
reset_num_timesteps: bool = True, reset_num_timesteps: bool = True,
) -> OffPolicyAlgorithm: ) -> OffPolicyAlgorithm:
return super(DDPG, self).learn( return super().learn(
total_timesteps=total_timesteps, total_timesteps=total_timesteps,
callback=callback, callback=callback,
log_interval=log_interval, log_interval=log_interval,

View file

@ -95,7 +95,7 @@ class DQN(OffPolicyAlgorithm):
_init_setup_model: bool = True, _init_setup_model: bool = True,
): ):
super(DQN, self).__init__( super().__init__(
policy, policy,
env, env,
learning_rate, learning_rate,
@ -138,7 +138,7 @@ class DQN(OffPolicyAlgorithm):
self._setup_model() self._setup_model()
def _setup_model(self) -> None: def _setup_model(self) -> None:
super(DQN, self)._setup_model() super()._setup_model()
self._create_aliases() self._create_aliases()
self.exploration_schedule = get_linear_fn( self.exploration_schedule = get_linear_fn(
self.exploration_initial_eps, self.exploration_initial_eps,
@ -261,7 +261,7 @@ class DQN(OffPolicyAlgorithm):
reset_num_timesteps: bool = True, reset_num_timesteps: bool = True,
) -> OffPolicyAlgorithm: ) -> OffPolicyAlgorithm:
return super(DQN, self).learn( return super().learn(
total_timesteps=total_timesteps, total_timesteps=total_timesteps,
callback=callback, callback=callback,
log_interval=log_interval, log_interval=log_interval,
@ -274,7 +274,7 @@ class DQN(OffPolicyAlgorithm):
) )
def _excluded_save_params(self) -> List[str]: def _excluded_save_params(self) -> List[str]:
return super(DQN, self)._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"]

View file

@ -37,7 +37,7 @@ class QNetwork(BasePolicy):
activation_fn: Type[nn.Module] = nn.ReLU, activation_fn: Type[nn.Module] = nn.ReLU,
normalize_images: bool = True, normalize_images: bool = True,
): ):
super(QNetwork, self).__init__( super().__init__(
observation_space, observation_space,
action_space, action_space,
features_extractor=features_extractor, features_extractor=features_extractor,
@ -118,7 +118,7 @@ class DQNPolicy(BasePolicy):
optimizer_class: Type[th.optim.Optimizer] = th.optim.Adam, optimizer_class: Type[th.optim.Optimizer] = th.optim.Adam,
optimizer_kwargs: Optional[Dict[str, Any]] = None, optimizer_kwargs: Optional[Dict[str, Any]] = None,
): ):
super(DQNPolicy, self).__init__( super().__init__(
observation_space, observation_space,
action_space, action_space,
features_extractor_class, features_extractor_class,
@ -239,7 +239,7 @@ class CnnPolicy(DQNPolicy):
optimizer_class: Type[th.optim.Optimizer] = th.optim.Adam, optimizer_class: Type[th.optim.Optimizer] = th.optim.Adam,
optimizer_kwargs: Optional[Dict[str, Any]] = None, optimizer_kwargs: Optional[Dict[str, Any]] = None,
): ):
super(CnnPolicy, self).__init__( super().__init__(
observation_space, observation_space,
action_space, action_space,
lr_schedule, lr_schedule,
@ -284,7 +284,7 @@ class MultiInputPolicy(DQNPolicy):
optimizer_class: Type[th.optim.Optimizer] = th.optim.Adam, optimizer_class: Type[th.optim.Optimizer] = th.optim.Adam,
optimizer_kwargs: Optional[Dict[str, Any]] = None, optimizer_kwargs: Optional[Dict[str, Any]] = None,
): ):
super(MultiInputPolicy, self).__init__( super().__init__(
observation_space, observation_space,
action_space, action_space,
lr_schedule, lr_schedule,

View file

@ -82,7 +82,7 @@ class HerReplayBuffer(DictReplayBuffer):
handle_timeout_termination: bool = True, handle_timeout_termination: bool = True,
): ):
super(HerReplayBuffer, self).__init__(buffer_size, env.observation_space, env.action_space, device, env.num_envs) super().__init__(buffer_size, env.observation_space, env.action_space, device, env.num_envs)
# convert goal_selection_strategy into GoalSelectionStrategy if string # convert goal_selection_strategy into GoalSelectionStrategy if string
if isinstance(goal_selection_strategy, str): if isinstance(goal_selection_strategy, str):

View file

@ -99,7 +99,7 @@ class PPO(OnPolicyAlgorithm):
_init_setup_model: bool = True, _init_setup_model: bool = True,
): ):
super(PPO, self).__init__( super().__init__(
policy, policy,
env, env,
learning_rate=learning_rate, learning_rate=learning_rate,
@ -162,7 +162,7 @@ class PPO(OnPolicyAlgorithm):
self._setup_model() self._setup_model()
def _setup_model(self) -> None: def _setup_model(self) -> None:
super(PPO, self)._setup_model() super()._setup_model()
# Initialize schedules for policy/value clipping # Initialize schedules for policy/value clipping
self.clip_range = get_schedule_fn(self.clip_range) self.clip_range = get_schedule_fn(self.clip_range)
@ -307,7 +307,7 @@ class PPO(OnPolicyAlgorithm):
reset_num_timesteps: bool = True, reset_num_timesteps: bool = True,
) -> "PPO": ) -> "PPO":
return super(PPO, self).learn( return super().learn(
total_timesteps=total_timesteps, total_timesteps=total_timesteps,
callback=callback, callback=callback,
log_interval=log_interval, log_interval=log_interval,

View file

@ -65,7 +65,7 @@ class Actor(BasePolicy):
clip_mean: float = 2.0, clip_mean: float = 2.0,
normalize_images: bool = True, normalize_images: bool = True,
): ):
super(Actor, self).__init__( super().__init__(
observation_space, observation_space,
action_space, action_space,
features_extractor=features_extractor, features_extractor=features_extractor,
@ -237,7 +237,7 @@ class SACPolicy(BasePolicy):
n_critics: int = 2, n_critics: int = 2,
share_features_extractor: bool = True, share_features_extractor: bool = True,
): ):
super(SACPolicy, self).__init__( super().__init__(
observation_space, observation_space,
action_space, action_space,
features_extractor_class, features_extractor_class,
@ -424,7 +424,7 @@ class CnnPolicy(SACPolicy):
n_critics: int = 2, n_critics: int = 2,
share_features_extractor: bool = True, share_features_extractor: bool = True,
): ):
super(CnnPolicy, self).__init__( super().__init__(
observation_space, observation_space,
action_space, action_space,
lr_schedule, lr_schedule,
@ -495,7 +495,7 @@ class MultiInputPolicy(SACPolicy):
n_critics: int = 2, n_critics: int = 2,
share_features_extractor: bool = True, share_features_extractor: bool = True,
): ):
super(MultiInputPolicy, self).__init__( super().__init__(
observation_space, observation_space,
action_space, action_space,
lr_schedule, lr_schedule,

View file

@ -110,7 +110,7 @@ class SAC(OffPolicyAlgorithm):
_init_setup_model: bool = True, _init_setup_model: bool = True,
): ):
super(SAC, self).__init__( super().__init__(
policy, policy,
env, env,
learning_rate, learning_rate,
@ -150,7 +150,7 @@ class SAC(OffPolicyAlgorithm):
self._setup_model() self._setup_model()
def _setup_model(self) -> None: def _setup_model(self) -> None:
super(SAC, self)._setup_model() super()._setup_model()
self._create_aliases() self._create_aliases()
# Target entropy is used when learning the entropy coefficient # Target entropy is used when learning the entropy coefficient
if self.target_entropy == "auto": if self.target_entropy == "auto":
@ -248,7 +248,7 @@ class SAC(OffPolicyAlgorithm):
current_q_values = self.critic(replay_data.observations, replay_data.actions) current_q_values = self.critic(replay_data.observations, replay_data.actions)
# Compute critic loss # Compute critic loss
critic_loss = 0.5 * sum([F.mse_loss(current_q, target_q_values) for current_q in current_q_values]) critic_loss = 0.5 * sum(F.mse_loss(current_q, target_q_values) for current_q in current_q_values)
critic_losses.append(critic_loss.item()) critic_losses.append(critic_loss.item())
# Optimize the critic # Optimize the critic
@ -295,7 +295,7 @@ class SAC(OffPolicyAlgorithm):
reset_num_timesteps: bool = True, reset_num_timesteps: bool = True,
) -> OffPolicyAlgorithm: ) -> OffPolicyAlgorithm:
return super(SAC, self).learn( return super().learn(
total_timesteps=total_timesteps, total_timesteps=total_timesteps,
callback=callback, callback=callback,
log_interval=log_interval, log_interval=log_interval,
@ -308,7 +308,7 @@ class SAC(OffPolicyAlgorithm):
) )
def _excluded_save_params(self) -> List[str]: def _excluded_save_params(self) -> List[str]:
return super(SAC, self)._excluded_save_params() + ["actor", "critic", "critic_target"] return super()._excluded_save_params() + ["actor", "critic", "critic_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", "actor.optimizer", "critic.optimizer"] state_dicts = ["policy", "actor.optimizer", "critic.optimizer"]

View file

@ -42,7 +42,7 @@ class Actor(BasePolicy):
activation_fn: Type[nn.Module] = nn.ReLU, activation_fn: Type[nn.Module] = nn.ReLU,
normalize_images: bool = True, normalize_images: bool = True,
): ):
super(Actor, self).__init__( super().__init__(
observation_space, observation_space,
action_space, action_space,
features_extractor=features_extractor, features_extractor=features_extractor,
@ -121,7 +121,7 @@ class TD3Policy(BasePolicy):
n_critics: int = 2, n_critics: int = 2,
share_features_extractor: bool = True, share_features_extractor: bool = True,
): ):
super(TD3Policy, self).__init__( super().__init__(
observation_space, observation_space,
action_space, action_space,
features_extractor_class, features_extractor_class,
@ -283,7 +283,7 @@ class CnnPolicy(TD3Policy):
n_critics: int = 2, n_critics: int = 2,
share_features_extractor: bool = True, share_features_extractor: bool = True,
): ):
super(CnnPolicy, self).__init__( super().__init__(
observation_space, observation_space,
action_space, action_space,
lr_schedule, lr_schedule,
@ -337,7 +337,7 @@ class MultiInputPolicy(TD3Policy):
n_critics: int = 2, n_critics: int = 2,
share_features_extractor: bool = True, share_features_extractor: bool = True,
): ):
super(MultiInputPolicy, self).__init__( super().__init__(
observation_space, observation_space,
action_space, action_space,
lr_schedule, lr_schedule,

View file

@ -95,7 +95,7 @@ class TD3(OffPolicyAlgorithm):
_init_setup_model: bool = True, _init_setup_model: bool = True,
): ):
super(TD3, self).__init__( super().__init__(
policy, policy,
env, env,
learning_rate, learning_rate,
@ -129,7 +129,7 @@ class TD3(OffPolicyAlgorithm):
self._setup_model() self._setup_model()
def _setup_model(self) -> None: def _setup_model(self) -> None:
super(TD3, self)._setup_model() super()._setup_model()
self._create_aliases() self._create_aliases()
def _create_aliases(self) -> None: def _create_aliases(self) -> None:
@ -168,7 +168,7 @@ class TD3(OffPolicyAlgorithm):
current_q_values = self.critic(replay_data.observations, replay_data.actions) current_q_values = self.critic(replay_data.observations, replay_data.actions)
# Compute critic loss # Compute critic loss
critic_loss = sum([F.mse_loss(current_q, target_q_values) for current_q in current_q_values]) critic_loss = sum(F.mse_loss(current_q, target_q_values) for current_q in current_q_values)
critic_losses.append(critic_loss.item()) critic_losses.append(critic_loss.item())
# Optimize the critics # Optimize the critics
@ -208,7 +208,7 @@ class TD3(OffPolicyAlgorithm):
reset_num_timesteps: bool = True, reset_num_timesteps: bool = True,
) -> OffPolicyAlgorithm: ) -> OffPolicyAlgorithm:
return super(TD3, self).learn( return super().learn(
total_timesteps=total_timesteps, total_timesteps=total_timesteps,
callback=callback, callback=callback,
log_interval=log_interval, log_interval=log_interval,
@ -221,7 +221,7 @@ class TD3(OffPolicyAlgorithm):
) )
def _excluded_save_params(self) -> List[str]: def _excluded_save_params(self) -> List[str]:
return super(TD3, self)._excluded_save_params() + ["actor", "critic", "actor_target", "critic_target"] return super()._excluded_save_params() + ["actor", "critic", "actor_target", "critic_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", "actor.optimizer", "critic.optimizer"] state_dicts = ["policy", "actor.optimizer", "critic.optimizer"]

View file

@ -1 +1 @@
1.5.1a4 1.5.1a5

View file

@ -10,7 +10,7 @@ from stable_baselines3.common.policies import ActorCriticPolicy
class CustomEnv(gym.Env): class CustomEnv(gym.Env):
def __init__(self, max_steps=8): def __init__(self, max_steps=8):
super(CustomEnv, self).__init__() super().__init__()
self.observation_space = gym.spaces.Box(low=-1, high=1, shape=(2,), dtype=np.float32) self.observation_space = gym.spaces.Box(low=-1, high=1, shape=(2,), dtype=np.float32)
self.action_space = gym.spaces.Box(low=-1, high=1, shape=(2,), dtype=np.float32) self.action_space = gym.spaces.Box(low=-1, high=1, shape=(2,), dtype=np.float32)
self.max_steps = max_steps self.max_steps = max_steps
@ -54,7 +54,7 @@ class InfiniteHorizonEnv(gym.Env):
class CheckGAECallback(BaseCallback): class CheckGAECallback(BaseCallback):
def __init__(self): def __init__(self):
super(CheckGAECallback, self).__init__(verbose=0) super().__init__(verbose=0)
def _on_rollout_end(self): def _on_rollout_end(self):
buffer = self.model.rollout_buffer buffer = self.model.rollout_buffer
@ -99,7 +99,7 @@ class CustomPolicy(ActorCriticPolicy):
"""Custom Policy with a constant value function""" """Custom Policy with a constant value function"""
def __init__(self, *args, **kwargs): def __init__(self, *args, **kwargs):
super(CustomPolicy, self).__init__(*args, **kwargs) super().__init__(*args, **kwargs)
self.constant_value = 0.0 self.constant_value = 0.0
def forward(self, obs, deterministic=False): def forward(self, obs, deterministic=False):

View file

@ -156,7 +156,7 @@ def test_save_load(tmp_path, model_class, use_sde, online_sampling):
params = deepcopy(model.policy.state_dict()) params = deepcopy(model.policy.state_dict())
# Modify all parameters to be random values # Modify all parameters to be random values
random_params = dict((param_name, th.rand_like(param)) for param_name, param in params.items()) random_params = {param_name: th.rand_like(param) for param_name, param in params.items()}
# Update model parameters with the new random values # Update model parameters with the new random values
model.policy.load_state_dict(random_params) model.policy.load_state_dict(random_params)

View file

@ -14,7 +14,7 @@ def test_monitor(tmp_path):
""" """
env = gym.make("CartPole-v1") env = gym.make("CartPole-v1")
env.seed(0) env.seed(0)
monitor_file = os.path.join(str(tmp_path), "stable_baselines-test-{}.monitor.csv".format(uuid.uuid4())) monitor_file = os.path.join(str(tmp_path), f"stable_baselines-test-{uuid.uuid4()}.monitor.csv")
monitor_env = Monitor(env, monitor_file) monitor_env = Monitor(env, monitor_file)
monitor_env.reset() monitor_env.reset()
total_steps = 1000 total_steps = 1000
@ -37,7 +37,7 @@ def test_monitor(tmp_path):
assert sum(monitor_env.get_episode_rewards()) == sum(ep_rewards) assert sum(monitor_env.get_episode_rewards()) == sum(ep_rewards)
_ = monitor_env.get_episode_times() _ = monitor_env.get_episode_times()
with open(monitor_file, "rt") as file_handler: with open(monitor_file) as file_handler:
first_line = file_handler.readline() first_line = file_handler.readline()
assert first_line.startswith("#") assert first_line.startswith("#")
metadata = json.loads(first_line[1:]) metadata = json.loads(first_line[1:])
@ -56,7 +56,7 @@ def test_monitor_load_results(tmp_path):
tmp_path = str(tmp_path) tmp_path = str(tmp_path)
env1 = gym.make("CartPole-v1") env1 = gym.make("CartPole-v1")
env1.seed(0) env1.seed(0)
monitor_file1 = os.path.join(tmp_path, "stable_baselines-test-{}.monitor.csv".format(uuid.uuid4())) monitor_file1 = os.path.join(tmp_path, f"stable_baselines-test-{uuid.uuid4()}.monitor.csv")
monitor_env1 = Monitor(env1, monitor_file1) monitor_env1 = Monitor(env1, monitor_file1)
monitor_files = get_monitor_files(tmp_path) monitor_files = get_monitor_files(tmp_path)
@ -76,7 +76,7 @@ def test_monitor_load_results(tmp_path):
env2 = gym.make("CartPole-v1") env2 = gym.make("CartPole-v1")
env2.seed(0) env2.seed(0)
monitor_file2 = os.path.join(tmp_path, "stable_baselines-test-{}.monitor.csv".format(uuid.uuid4())) monitor_file2 = os.path.join(tmp_path, f"stable_baselines-test-{uuid.uuid4()}.monitor.csv")
monitor_env2 = Monitor(env2, monitor_file2) monitor_env2 = Monitor(env2, monitor_file2)
monitor_files = get_monitor_files(tmp_path) monitor_files = get_monitor_files(tmp_path)
assert len(monitor_files) == 2 assert len(monitor_files) == 2

View file

@ -64,7 +64,7 @@ def test_save_load(tmp_path, model_class):
model.set_parameters(invalid_object_params, exact_match=False) model.set_parameters(invalid_object_params, exact_match=False)
# Test that exact_match catches when something was missed. # Test that exact_match catches when something was missed.
missing_object_params = dict((k, v) for k, v in list(original_params.items())[:-1]) missing_object_params = {k: v for k, v in list(original_params.items())[:-1]}
with pytest.raises(ValueError): with pytest.raises(ValueError):
model.set_parameters(missing_object_params, exact_match=True) model.set_parameters(missing_object_params, exact_match=True)
@ -446,7 +446,7 @@ def test_save_load_policy(tmp_path, model_class, policy_str, use_sde):
params = deepcopy(policy.state_dict()) params = deepcopy(policy.state_dict())
# Modify all parameters to be random values # Modify all parameters to be random values
random_params = dict((param_name, th.rand_like(param)) for param_name, param in params.items()) random_params = {param_name: th.rand_like(param) for param_name, param in params.items()}
# Update model parameters with the new random values # Update model parameters with the new random values
policy.load_state_dict(random_params) policy.load_state_dict(random_params)
@ -537,7 +537,7 @@ def test_save_load_q_net(tmp_path, model_class, policy_str):
params = deepcopy(q_net.state_dict()) params = deepcopy(q_net.state_dict())
# Modify all parameters to be random values # Modify all parameters to be random values
random_params = dict((param_name, th.rand_like(param)) for param_name, param in params.items()) random_params = {param_name: th.rand_like(param) for param_name, param in params.items()}
# Update model parameters with the new random values # Update model parameters with the new random values
q_net.load_state_dict(random_params) q_net.load_state_dict(random_params)

View file

@ -9,7 +9,7 @@ from stable_baselines3.common.evaluation import evaluate_policy
class DummyMultiDiscreteSpace(gym.Env): class DummyMultiDiscreteSpace(gym.Env):
def __init__(self, nvec): def __init__(self, nvec):
super(DummyMultiDiscreteSpace, self).__init__() super().__init__()
self.observation_space = gym.spaces.MultiDiscrete(nvec) self.observation_space = gym.spaces.MultiDiscrete(nvec)
self.action_space = gym.spaces.Box(low=-1, high=1, shape=(2,), dtype=np.float32) self.action_space = gym.spaces.Box(low=-1, high=1, shape=(2,), dtype=np.float32)
@ -22,7 +22,7 @@ class DummyMultiDiscreteSpace(gym.Env):
class DummyMultiBinary(gym.Env): class DummyMultiBinary(gym.Env):
def __init__(self, n): def __init__(self, n):
super(DummyMultiBinary, self).__init__() super().__init__()
self.observation_space = gym.spaces.MultiBinary(n) self.observation_space = gym.spaces.MultiBinary(n)
self.action_space = gym.spaces.Box(low=-1, high=1, shape=(2,), dtype=np.float32) self.action_space = gym.spaces.Box(low=-1, high=1, shape=(2,), dtype=np.float32)

View file

@ -28,7 +28,7 @@ class FlattenBatchNormDropoutExtractor(BaseFeaturesExtractor):
""" """
def __init__(self, observation_space: gym.Space): def __init__(self, observation_space: gym.Space):
super(FlattenBatchNormDropoutExtractor, self).__init__( super().__init__(
observation_space, observation_space,
get_flattened_obs_dim(observation_space), get_flattened_obs_dim(observation_space),
) )

View file

@ -180,7 +180,7 @@ class AlwaysDoneWrapper(gym.Wrapper):
# Pretends that environment only has single step for each # Pretends that environment only has single step for each
# episode. # episode.
def __init__(self, env): def __init__(self, env):
super(AlwaysDoneWrapper, self).__init__(env) super().__init__(env)
self.last_obs = None self.last_obs = None
self.needs_reset = True self.needs_reset = True

View file

@ -12,7 +12,7 @@ class NanAndInfEnv(gym.Env):
metadata = {"render.modes": ["human"]} metadata = {"render.modes": ["human"]}
def __init__(self): def __init__(self):
super(NanAndInfEnv, self).__init__() super().__init__()
self.action_space = spaces.Box(low=-np.inf, high=np.inf, shape=(1,), dtype=np.float64) self.action_space = spaces.Box(low=-np.inf, high=np.inf, shape=(1,), dtype=np.float64)
self.observation_space = spaces.Box(low=-np.inf, high=np.inf, shape=(1,), dtype=np.float64) self.observation_space = spaces.Box(low=-np.inf, high=np.inf, shape=(1,), dtype=np.float64)

View file

@ -36,7 +36,7 @@ def test_vec_monitor(tmp_path):
monitor_env.close() monitor_env.close()
with open(monitor_file, "rt") as file_handler: with open(monitor_file) as file_handler:
first_line = file_handler.readline() first_line = file_handler.readline()
assert first_line.startswith("#") assert first_line.startswith("#")
metadata = json.loads(first_line[1:]) metadata = json.loads(first_line[1:])
@ -66,7 +66,7 @@ def test_vec_monitor_info_keywords(tmp_path):
monitor_env.close() monitor_env.close()
with open(monitor_file, "rt") as f: with open(monitor_file) as f:
reader = csv.reader(f) reader = csv.reader(f)
for i, line in enumerate(reader): for i, line in enumerate(reader):
if i == 0 or i == 1: if i == 0 or i == 1:

View file

@ -47,7 +47,7 @@ class DummyDictEnv(gym.GoalEnv):
""" """
def __init__(self): def __init__(self):
super(DummyDictEnv, self).__init__() super().__init__()
self.observation_space = spaces.Dict( self.observation_space = spaces.Dict(
{ {
"observation": spaces.Box(low=-20.0, high=20.0, shape=(4,), dtype=np.float32), "observation": spaces.Box(low=-20.0, high=20.0, shape=(4,), dtype=np.float32),