mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-09-04 20:23:54 +00:00
Add custom arch for off-policy actor/critic networks (#182)
* Add custom arch for off-policy actor/critic networks * Fix type hints * Address comments * Make sure number of updated parameters match in polyak * Add zip_strict for strict-length zipping * Fix building docs * Add test for zip strict * Faster tests Co-authored-by: Anssi "Miffyli" Kanervisto <kaneran21@hotmail.com>
This commit is contained in:
parent
fc9527157a
commit
2599f04940
11 changed files with 177 additions and 39 deletions
|
|
@ -258,9 +258,31 @@ If your task requires even more granular control over the policy/value architect
|
|||
|
||||
|
||||
|
||||
.. TODO (see https://github.com/DLR-RM/stable-baselines3/issues/113)
|
||||
.. Off-Policy Algorithms
|
||||
.. ^^^^^^^^^^^^^^^^^^^^^
|
||||
..
|
||||
.. If you need a network architecture that is different for the actor and the critic when using ``SAC``, ``DDPG`` or ``TD3``,
|
||||
.. you can easily redefine the actor class for instance.
|
||||
Off-Policy Algorithms
|
||||
^^^^^^^^^^^^^^^^^^^^^
|
||||
|
||||
If you need a network architecture that is different for the actor and the critic when using ``SAC``, ``DDPG`` or ``TD3``,
|
||||
you can pass a dictionary of the following structure: ``dict(qf=[<critic network architecture>], pi=[<actor network architecture>])``.
|
||||
|
||||
For example, if you want a different architecture for the actor (aka ``pi``) and the critic (Q-function aka ``qf``) networks,
|
||||
then you can specify ``net_arch=dict(qf=[400, 300], pi=[64, 64])``.
|
||||
|
||||
Otherwise, to have actor and critic that share the same network architecture,
|
||||
you only need to specify ``net_arch=[256, 256]`` (here, two hidden layers of 256 units each).
|
||||
|
||||
|
||||
.. note::
|
||||
Compared to their on-policy counterparts, no shared layers (other than the feature extractor)
|
||||
between the actor and the critic are allowed (to prevent issues with target networks).
|
||||
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
from stable_baselines3 import SAC
|
||||
|
||||
# Custom actor architecture with two layers of 64 units each
|
||||
# Custom critic architecture with two layers of 400 and 300 units
|
||||
policy_kwargs = dict(net_arch=dict(pi=[64, 64], qf=[400, 300]))
|
||||
# Create the agent
|
||||
model = SAC("MlpPolicy", "Pendulum-v0", policy_kwargs=policy_kwargs, verbose=1)
|
||||
model.learn(5000)
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ Breaking Changes:
|
|||
|
||||
New Features:
|
||||
^^^^^^^^^^^^^
|
||||
- Allow custom actor/critic network architectures using ``net_arch=dict(qf=[400, 300], pi=[64, 64])`` for off-policy algorithms (SAC, TD3, DDPG)
|
||||
|
||||
Bug Fixes:
|
||||
^^^^^^^^^^
|
||||
|
|
|
|||
|
|
@ -219,3 +219,43 @@ class MlpExtractor(nn.Module):
|
|||
"""
|
||||
shared_latent = self.shared_net(features)
|
||||
return self.policy_net(shared_latent), self.value_net(shared_latent)
|
||||
|
||||
|
||||
def get_actor_critic_arch(net_arch: Union[List[int], Dict[str, List[int]]]) -> Tuple[List[int], List[int]]:
|
||||
"""
|
||||
Get the actor and critic network architectures for off-policy actor-critic algorithms (SAC, TD3, DDPG).
|
||||
|
||||
The ``net_arch`` parameter allows to specify the amount and size of the hidden layers,
|
||||
which can be different for the actor and the critic.
|
||||
It is assumed to be a list of ints or a dict.
|
||||
|
||||
1. If it is a list, actor and critic networks will have the same architecture.
|
||||
The architecture is represented by a list of integers (of arbitrary length (zero allowed))
|
||||
each specifying the number of units per layer.
|
||||
If the number of ints is zero, the network will be linear.
|
||||
2. If it is a dict, it should have the following structure:
|
||||
``dict(qf=[<critic network architecture>], pi=[<actor network architecture>])``.
|
||||
where the network architecture is a list as described in 1.
|
||||
|
||||
For example, to have actor and critic that share the same network architecture,
|
||||
you only need to specify ``net_arch=[256, 256]`` (here, two hidden layers of 256 units each).
|
||||
|
||||
If you want a different architecture for the actor and the critic,
|
||||
then you can specify ``net_arch=dict(qf=[400, 300], pi=[64, 64])``.
|
||||
|
||||
.. note::
|
||||
Compared to their on-policy counterparts, no shared layers (other than the feature extractor)
|
||||
between the actor and the critic are allowed (to prevent issues with target networks).
|
||||
|
||||
:param net_arch: The specification of the actor and critic networks.
|
||||
See above for details on its formatting.
|
||||
:return: The network architectures for the actor and the critic
|
||||
"""
|
||||
if isinstance(net_arch, list):
|
||||
actor_arch, critic_arch = net_arch, net_arch
|
||||
else:
|
||||
assert isinstance(net_arch, dict), "Error: the net_arch can only contain be a list of ints or a dict"
|
||||
assert "pi" in net_arch, "Error: no key 'pi' was provided in net_arch for the actor network"
|
||||
assert "qf" in net_arch, "Error: no key 'qf' was provided in net_arch for the critic network"
|
||||
actor_arch, critic_arch = net_arch["pi"], net_arch["qf"]
|
||||
return actor_arch, critic_arch
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import glob
|
|||
import os
|
||||
import random
|
||||
from collections import deque
|
||||
from itertools import zip_longest
|
||||
from typing import Callable, Iterable, Optional, Union
|
||||
|
||||
import gym
|
||||
|
|
@ -286,6 +287,24 @@ def safe_mean(arr: Union[np.ndarray, list, deque]) -> np.ndarray:
|
|||
return np.nan if len(arr) == 0 else np.mean(arr)
|
||||
|
||||
|
||||
def zip_strict(*iterables: Iterable) -> Iterable:
|
||||
r"""
|
||||
``zip()`` function but enforces that iterables are of equal length.
|
||||
Raises ``ValueError`` if iterables not of equal length.
|
||||
Code inspired by Stackoverflow answer for question #32954486.
|
||||
|
||||
:param \*iterables: iterables to ``zip()``
|
||||
"""
|
||||
# As in Stackoverflow #32954486, use
|
||||
# new object for "empty" in case we have
|
||||
# Nones in iterable.
|
||||
sentinel = object()
|
||||
for combo in zip_longest(*iterables, fillvalue=sentinel):
|
||||
if sentinel in combo:
|
||||
raise ValueError("Iterables have different lengths")
|
||||
yield combo
|
||||
|
||||
|
||||
def polyak_update(params: Iterable[th.nn.Parameter], target_params: Iterable[th.nn.Parameter], tau: float) -> None:
|
||||
"""
|
||||
Perform a Polyak average update on ``target_params`` using ``params``:
|
||||
|
|
@ -303,6 +322,7 @@ def polyak_update(params: Iterable[th.nn.Parameter], target_params: Iterable[th.
|
|||
:param tau: the soft update coefficient ("Polyak update", between 0 and 1)
|
||||
"""
|
||||
with th.no_grad():
|
||||
for param, target_param in zip(params, target_params):
|
||||
# zip does not raise an exception if length of parameters does not match.
|
||||
for param, target_param in zip_strict(params, target_params):
|
||||
target_param.data.mul_(1 - tau)
|
||||
th.add(target_param.data, param.data, alpha=tau, out=target_param.data)
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import Any, Callable, Dict, List, Optional, Tuple, Type
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple, Type, Union
|
||||
|
||||
import gym
|
||||
import torch as th
|
||||
|
|
@ -7,7 +7,13 @@ from torch import nn
|
|||
from stable_baselines3.common.distributions import SquashedDiagGaussianDistribution, StateDependentNoiseDistribution
|
||||
from stable_baselines3.common.policies import BasePolicy, ContinuousCritic, create_sde_features_extractor, register_policy
|
||||
from stable_baselines3.common.preprocessing import get_action_dim
|
||||
from stable_baselines3.common.torch_layers import BaseFeaturesExtractor, FlattenExtractor, NatureCNN, create_mlp
|
||||
from stable_baselines3.common.torch_layers import (
|
||||
BaseFeaturesExtractor,
|
||||
FlattenExtractor,
|
||||
NatureCNN,
|
||||
create_mlp,
|
||||
get_actor_critic_arch,
|
||||
)
|
||||
|
||||
# CAP the standard deviation of the actor
|
||||
LOG_STD_MAX = 2
|
||||
|
|
@ -220,7 +226,7 @@ class SACPolicy(BasePolicy):
|
|||
observation_space: gym.spaces.Space,
|
||||
action_space: gym.spaces.Space,
|
||||
lr_schedule: Callable,
|
||||
net_arch: Optional[List[int]] = None,
|
||||
net_arch: Optional[Union[List[int], Dict[str, List[int]]]] = None,
|
||||
activation_fn: Type[nn.Module] = nn.ReLU,
|
||||
use_sde: bool = False,
|
||||
log_std_init: float = -3,
|
||||
|
|
@ -250,6 +256,8 @@ class SACPolicy(BasePolicy):
|
|||
else:
|
||||
net_arch = []
|
||||
|
||||
actor_arch, critic_arch = get_actor_critic_arch(net_arch)
|
||||
|
||||
# Create shared features extractor
|
||||
self.features_extractor = features_extractor_class(self.observation_space, **self.features_extractor_kwargs)
|
||||
self.features_dim = self.features_extractor.features_dim
|
||||
|
|
@ -261,7 +269,7 @@ class SACPolicy(BasePolicy):
|
|||
"action_space": self.action_space,
|
||||
"features_extractor": self.features_extractor,
|
||||
"features_dim": self.features_dim,
|
||||
"net_arch": self.net_arch,
|
||||
"net_arch": actor_arch,
|
||||
"activation_fn": self.activation_fn,
|
||||
"normalize_images": normalize_images,
|
||||
}
|
||||
|
|
@ -275,7 +283,7 @@ class SACPolicy(BasePolicy):
|
|||
}
|
||||
self.actor_kwargs.update(sde_kwargs)
|
||||
self.critic_kwargs = self.net_args.copy()
|
||||
self.critic_kwargs.update({"n_critics": n_critics})
|
||||
self.critic_kwargs.update({"n_critics": n_critics, "net_arch": critic_arch})
|
||||
|
||||
self.actor, self.actor_target = None, None
|
||||
self.critic, self.critic_target = None, None
|
||||
|
|
@ -300,7 +308,7 @@ class SACPolicy(BasePolicy):
|
|||
|
||||
data.update(
|
||||
dict(
|
||||
net_arch=self.net_args["net_arch"],
|
||||
net_arch=self.net_arch,
|
||||
activation_fn=self.net_args["activation_fn"],
|
||||
use_sde=self.actor_kwargs["use_sde"],
|
||||
log_std_init=self.actor_kwargs["log_std_init"],
|
||||
|
|
@ -374,7 +382,7 @@ class CnnPolicy(SACPolicy):
|
|||
observation_space: gym.spaces.Space,
|
||||
action_space: gym.spaces.Space,
|
||||
lr_schedule: Callable,
|
||||
net_arch: Optional[List[int]] = None,
|
||||
net_arch: Optional[Union[List[int], Dict[str, List[int]]]] = None,
|
||||
activation_fn: Type[nn.Module] = nn.ReLU,
|
||||
use_sde: bool = False,
|
||||
log_std_init: float = -3,
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import Any, Callable, Dict, List, Optional, Type
|
||||
from typing import Any, Callable, Dict, List, Optional, Type, Union
|
||||
|
||||
import gym
|
||||
import torch as th
|
||||
|
|
@ -6,7 +6,13 @@ from torch import nn
|
|||
|
||||
from stable_baselines3.common.policies import BasePolicy, ContinuousCritic, register_policy
|
||||
from stable_baselines3.common.preprocessing import get_action_dim
|
||||
from stable_baselines3.common.torch_layers import BaseFeaturesExtractor, FlattenExtractor, NatureCNN, create_mlp
|
||||
from stable_baselines3.common.torch_layers import (
|
||||
BaseFeaturesExtractor,
|
||||
FlattenExtractor,
|
||||
NatureCNN,
|
||||
create_mlp,
|
||||
get_actor_critic_arch,
|
||||
)
|
||||
|
||||
|
||||
class Actor(BasePolicy):
|
||||
|
|
@ -101,7 +107,7 @@ class TD3Policy(BasePolicy):
|
|||
observation_space: gym.spaces.Space,
|
||||
action_space: gym.spaces.Space,
|
||||
lr_schedule: Callable,
|
||||
net_arch: Optional[List[int]] = None,
|
||||
net_arch: Optional[Union[List[int], Dict[str, List[int]]]] = None,
|
||||
activation_fn: Type[nn.Module] = nn.ReLU,
|
||||
features_extractor_class: Type[BaseFeaturesExtractor] = FlattenExtractor,
|
||||
features_extractor_kwargs: Optional[Dict[str, Any]] = None,
|
||||
|
|
@ -127,6 +133,8 @@ class TD3Policy(BasePolicy):
|
|||
else:
|
||||
net_arch = []
|
||||
|
||||
actor_arch, critic_arch = get_actor_critic_arch(net_arch)
|
||||
|
||||
self.features_extractor = features_extractor_class(self.observation_space, **self.features_extractor_kwargs)
|
||||
self.features_dim = self.features_extractor.features_dim
|
||||
|
||||
|
|
@ -137,12 +145,13 @@ class TD3Policy(BasePolicy):
|
|||
"action_space": self.action_space,
|
||||
"features_extractor": self.features_extractor,
|
||||
"features_dim": self.features_dim,
|
||||
"net_arch": self.net_arch,
|
||||
"net_arch": actor_arch,
|
||||
"activation_fn": self.activation_fn,
|
||||
"normalize_images": normalize_images,
|
||||
}
|
||||
self.actor_kwargs = self.net_args.copy()
|
||||
self.critic_kwargs = self.net_args.copy()
|
||||
self.critic_kwargs.update({"n_critics": n_critics})
|
||||
self.critic_kwargs.update({"n_critics": n_critics, "net_arch": critic_arch})
|
||||
self.actor, self.actor_target = None, None
|
||||
self.critic, self.critic_target = None, None
|
||||
|
||||
|
|
@ -163,7 +172,7 @@ class TD3Policy(BasePolicy):
|
|||
|
||||
data.update(
|
||||
dict(
|
||||
net_arch=self.net_args["net_arch"],
|
||||
net_arch=self.net_arch,
|
||||
activation_fn=self.net_args["activation_fn"],
|
||||
n_critics=self.critic_kwargs["n_critics"],
|
||||
lr_schedule=self._dummy_schedule, # dummy lr schedule, not needed for loading policy alone
|
||||
|
|
@ -176,7 +185,7 @@ class TD3Policy(BasePolicy):
|
|||
return data
|
||||
|
||||
def make_actor(self) -> Actor:
|
||||
return Actor(**self.net_args).to(self.device)
|
||||
return Actor(**self.actor_kwargs).to(self.device)
|
||||
|
||||
def make_critic(self) -> ContinuousCritic:
|
||||
return ContinuousCritic(**self.critic_kwargs).to(self.device)
|
||||
|
|
@ -217,7 +226,7 @@ class CnnPolicy(TD3Policy):
|
|||
observation_space: gym.spaces.Space,
|
||||
action_space: gym.spaces.Space,
|
||||
lr_schedule: Callable,
|
||||
net_arch: Optional[List[int]] = None,
|
||||
net_arch: Optional[Union[List[int], Dict[str, List[int]]]] = None,
|
||||
activation_fn: Type[nn.Module] = nn.ReLU,
|
||||
features_extractor_class: Type[BaseFeaturesExtractor] = NatureCNN,
|
||||
features_extractor_kwargs: Optional[Dict[str, Any]] = None,
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import pytest
|
||||
import torch as th
|
||||
|
||||
from stable_baselines3 import A2C, PPO, SAC, TD3
|
||||
from stable_baselines3 import A2C, DQN, PPO, SAC, TD3
|
||||
from stable_baselines3.common.sb2_compat.rmsprop_tf_like import RMSpropTFLike
|
||||
|
||||
|
||||
|
|
@ -19,22 +19,33 @@ from stable_baselines3.common.sb2_compat.rmsprop_tf_like import RMSpropTFLike
|
|||
)
|
||||
@pytest.mark.parametrize("model_class", [A2C, PPO])
|
||||
def test_flexible_mlp(model_class, net_arch):
|
||||
_ = model_class("MlpPolicy", "CartPole-v1", policy_kwargs=dict(net_arch=net_arch), n_steps=100).learn(1000)
|
||||
_ = model_class("MlpPolicy", "CartPole-v1", policy_kwargs=dict(net_arch=net_arch), n_steps=100).learn(300)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("net_arch", [[4], [4, 4]])
|
||||
@pytest.mark.parametrize("net_arch", [[], [4], [4, 4], dict(qf=[8], pi=[8, 4])])
|
||||
@pytest.mark.parametrize("model_class", [SAC, TD3])
|
||||
def test_custom_offpolicy(model_class, net_arch):
|
||||
_ = model_class("MlpPolicy", "Pendulum-v0", policy_kwargs=dict(net_arch=net_arch)).learn(1000)
|
||||
_ = model_class("MlpPolicy", "Pendulum-v0", policy_kwargs=dict(net_arch=net_arch), learning_starts=100).learn(300)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_class", [A2C, PPO, SAC, TD3])
|
||||
@pytest.mark.parametrize("optimizer_kwargs", [None, dict(weight_decay=0.0)])
|
||||
def test_custom_optimizer(model_class, optimizer_kwargs):
|
||||
kwargs = {}
|
||||
if model_class in {DQN, SAC, TD3}:
|
||||
kwargs = dict(learning_starts=100)
|
||||
elif model_class in {A2C, PPO}:
|
||||
kwargs = dict(n_steps=100)
|
||||
|
||||
policy_kwargs = dict(optimizer_class=th.optim.AdamW, optimizer_kwargs=optimizer_kwargs, net_arch=[32])
|
||||
_ = model_class("MlpPolicy", "Pendulum-v0", policy_kwargs=policy_kwargs).learn(1000)
|
||||
_ = model_class("MlpPolicy", "Pendulum-v0", policy_kwargs=policy_kwargs, **kwargs).learn(300)
|
||||
|
||||
|
||||
def test_tf_like_rmsprop_optimizer():
|
||||
policy_kwargs = dict(optimizer_class=RMSpropTFLike, net_arch=[32])
|
||||
_ = A2C("MlpPolicy", "Pendulum-v0", policy_kwargs=policy_kwargs).learn(1000)
|
||||
_ = A2C("MlpPolicy", "Pendulum-v0", policy_kwargs=policy_kwargs).learn(500)
|
||||
|
||||
|
||||
def test_dqn_custom_policy():
|
||||
policy_kwargs = dict(optimizer_class=RMSpropTFLike, net_arch=[32])
|
||||
_ = DQN("MlpPolicy", "CartPole-v1", policy_kwargs=policy_kwargs, learning_starts=100).learn(300)
|
||||
|
|
|
|||
|
|
@ -80,7 +80,7 @@ def test_n_critics(n_critics):
|
|||
model = SAC(
|
||||
"MlpPolicy", "Pendulum-v0", policy_kwargs=dict(net_arch=[64, 64], n_critics=n_critics), learning_starts=100, verbose=1
|
||||
)
|
||||
model.learn(total_timesteps=1000)
|
||||
model.learn(total_timesteps=500)
|
||||
|
||||
|
||||
def test_dqn():
|
||||
|
|
@ -88,10 +88,10 @@ def test_dqn():
|
|||
"MlpPolicy",
|
||||
"CartPole-v1",
|
||||
policy_kwargs=dict(net_arch=[64, 64]),
|
||||
learning_starts=500,
|
||||
learning_starts=100,
|
||||
buffer_size=500,
|
||||
learning_rate=3e-4,
|
||||
verbose=1,
|
||||
create_eval_env=True,
|
||||
)
|
||||
model.learn(total_timesteps=1000, eval_freq=500)
|
||||
model.learn(total_timesteps=500, eval_freq=250)
|
||||
|
|
|
|||
|
|
@ -45,7 +45,7 @@ def test_save_load(tmp_path, model_class):
|
|||
|
||||
# create model
|
||||
model = model_class("MlpPolicy", env, policy_kwargs=dict(net_arch=[16]), verbose=1)
|
||||
model.learn(total_timesteps=500, eval_freq=250)
|
||||
model.learn(total_timesteps=500)
|
||||
|
||||
env.reset()
|
||||
observations = np.concatenate([env.step([env.action_space.sample()])[0] for _ in range(10)], axis=0)
|
||||
|
|
@ -154,7 +154,7 @@ def test_save_load(tmp_path, model_class):
|
|||
assert np.allclose(selected_actions, new_selected_actions, 1e-4)
|
||||
|
||||
# check if learn still works
|
||||
model.learn(total_timesteps=1000, eval_freq=500)
|
||||
model.learn(total_timesteps=500)
|
||||
|
||||
del model
|
||||
|
||||
|
|
@ -174,20 +174,26 @@ def test_set_env(model_class):
|
|||
env2 = DummyVecEnv([lambda: select_env(model_class)])
|
||||
env3 = select_env(model_class)
|
||||
|
||||
kwargs = {}
|
||||
if model_class in {DQN, DDPG, SAC, TD3}:
|
||||
kwargs = dict(learning_starts=100)
|
||||
elif model_class in {A2C, PPO}:
|
||||
kwargs = dict(n_steps=100)
|
||||
|
||||
# create model
|
||||
model = model_class("MlpPolicy", env, policy_kwargs=dict(net_arch=[16]))
|
||||
model = model_class("MlpPolicy", env, policy_kwargs=dict(net_arch=[16]), **kwargs)
|
||||
# learn
|
||||
model.learn(total_timesteps=1000, eval_freq=500)
|
||||
model.learn(total_timesteps=300)
|
||||
|
||||
# change env
|
||||
model.set_env(env2)
|
||||
# learn again
|
||||
model.learn(total_timesteps=1000, eval_freq=500)
|
||||
model.learn(total_timesteps=300)
|
||||
|
||||
# change env test wrapping
|
||||
model.set_env(env3)
|
||||
# learn again
|
||||
model.learn(total_timesteps=1000, eval_freq=500)
|
||||
model.learn(total_timesteps=300)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_class", MODEL_LIST)
|
||||
|
|
@ -309,7 +315,7 @@ def test_save_load_policy(tmp_path, model_class, policy_str):
|
|||
|
||||
# create model
|
||||
model = model_class(policy_str, env, policy_kwargs=dict(net_arch=[16]), verbose=1, **kwargs)
|
||||
model.learn(total_timesteps=500, eval_freq=250)
|
||||
model.learn(total_timesteps=500)
|
||||
|
||||
env.reset()
|
||||
observations = np.concatenate([env.step([env.action_space.sample()])[0] for _ in range(10)], axis=0)
|
||||
|
|
|
|||
|
|
@ -68,3 +68,6 @@ def test_state_dependent_offpolicy_noise(model_class, sde_net_arch, use_expln):
|
|||
policy_kwargs=dict(log_std_init=-2, sde_net_arch=sde_net_arch, use_expln=use_expln),
|
||||
)
|
||||
model.learn(total_timesteps=int(500), eval_freq=250)
|
||||
model.policy.reset_noise()
|
||||
if model_class == SAC:
|
||||
model.policy.actor.get_std()
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ from stable_baselines3.common.cmd_util import make_atari_env, make_vec_env
|
|||
from stable_baselines3.common.evaluation import evaluate_policy
|
||||
from stable_baselines3.common.monitor import Monitor
|
||||
from stable_baselines3.common.noise import ActionNoise, OrnsteinUhlenbeckActionNoise, VectorizedActionNoise
|
||||
from stable_baselines3.common.utils import polyak_update
|
||||
from stable_baselines3.common.utils import polyak_update, zip_strict
|
||||
from stable_baselines3.common.vec_env import DummyVecEnv, SubprocVecEnv
|
||||
|
||||
|
||||
|
|
@ -167,3 +167,21 @@ def test_polyak():
|
|||
|
||||
assert th.allclose(param1, target1)
|
||||
assert th.allclose(param2, target2)
|
||||
|
||||
|
||||
def test_zip_strict():
|
||||
# Iterables with different lengths
|
||||
list_a = [0, 1]
|
||||
list_b = [1, 2, 3]
|
||||
# zip does not raise any error
|
||||
for _, _ in zip(list_a, list_b):
|
||||
pass
|
||||
|
||||
# zip_strict does raise an error
|
||||
with pytest.raises(ValueError):
|
||||
for _, _ in zip_strict(list_a, list_b):
|
||||
pass
|
||||
|
||||
# same length, should not raise an error
|
||||
for _, _ in zip_strict(list_a, list_b[: len(list_a)]):
|
||||
pass
|
||||
|
|
|
|||
Loading…
Reference in a new issue