From 037986a91d272a79cb6a2d434fedf5fb27f23904 Mon Sep 17 00:00:00 2001 From: Antonin Raffin Date: Wed, 11 Mar 2020 16:35:13 +0100 Subject: [PATCH] Add test for `expln` --- docs/misc/changelog.rst | 1 + tests/test_sde.py | 11 ++++++----- 2 files changed, 7 insertions(+), 5 deletions(-) diff --git a/docs/misc/changelog.rst b/docs/misc/changelog.rst index 45d28de..3bd119b 100644 --- a/docs/misc/changelog.rst +++ b/docs/misc/changelog.rst @@ -27,6 +27,7 @@ Others: - Added ``clip_mean`` parameter to SAC policy - Buffers now return ``NamedTuple`` - More typing +- Add test for ``expln`` Documentation: ^^^^^^^^^^^^^^ diff --git a/tests/test_sde.py b/tests/test_sde.py index 497f398..96f6c8f 100644 --- a/tests/test_sde.py +++ b/tests/test_sde.py @@ -2,7 +2,7 @@ import pytest import torch as th from torch.distributions import Normal -from torchy_baselines import A2C, TD3, SAC +from torchy_baselines import A2C, TD3, SAC, PPO def test_state_dependent_exploration_grad(): @@ -55,12 +55,13 @@ def test_state_dependent_exploration_grad(): assert sigma_hat.grad.allclose(grad) -@pytest.mark.parametrize("model_class", [TD3, SAC, A2C]) +@pytest.mark.parametrize("model_class", [TD3, SAC, A2C, PPO]) @pytest.mark.parametrize("sde_net_arch", [None, [32, 16], []]) -def test_state_dependent_offpolicy_noise(model_class, sde_net_arch): +@pytest.mark.parametrize("use_expln", [False, True]) +def test_state_dependent_offpolicy_noise(model_class, sde_net_arch, use_expln): model = model_class('MlpPolicy', 'Pendulum-v0', use_sde=True, seed=None, create_eval_env=True, - verbose=1, policy_kwargs=dict(log_std_init=-2, sde_net_arch=sde_net_arch)) - model.learn(total_timesteps=int(1000), eval_freq=500) + verbose=1, 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) def test_scheduler():