From 9c77a6d9ec457a683cd5caebc8eca142517b7d1b Mon Sep 17 00:00:00 2001 From: Antonin RAFFIN Date: Wed, 5 Aug 2020 13:43:50 +0200 Subject: [PATCH] Pass gamma to n-step replay buffer --- stable_baselines3/common/off_policy_algorithm.py | 4 +++- tests/test_buffers.py | 2 +- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/stable_baselines3/common/off_policy_algorithm.py b/stable_baselines3/common/off_policy_algorithm.py index 36ad34b..a932c03 100644 --- a/stable_baselines3/common/off_policy_algorithm.py +++ b/stable_baselines3/common/off_policy_algorithm.py @@ -10,7 +10,7 @@ import torch as th from stable_baselines3.common import logger from stable_baselines3.common.base_class import BaseAlgorithm -from stable_baselines3.common.buffers import ReplayBuffer +from stable_baselines3.common.buffers import NstepReplayBuffer, ReplayBuffer from stable_baselines3.common.callbacks import BaseCallback from stable_baselines3.common.noise import ActionNoise from stable_baselines3.common.policies import BasePolicy @@ -150,6 +150,8 @@ class OffPolicyAlgorithm(BaseAlgorithm): self.replay_buffer = None self.replay_buffer_class = replay_buffer_class or ReplayBuffer self.replay_buffer_kwargs = dict(replay_buffer_kwargs or {}) + if self.replay_buffer_class == NstepReplayBuffer: + self.replay_buffer_kwargs["gamma"] = gamma def _setup_model(self) -> None: self._setup_lr_schedule() diff --git a/tests/test_buffers.py b/tests/test_buffers.py index 69d8d94..5d34346 100644 --- a/tests/test_buffers.py +++ b/tests/test_buffers.py @@ -108,7 +108,7 @@ def test_with_algo(algo): kwargs = { "policy_kwargs": dict(net_arch=[64]), "replay_buffer_class": NstepReplayBuffer, - "replay_buffer_kwargs": dict(gamma=0.99, n_step=10), + "replay_buffer_kwargs": dict(n_step=10), } if algo in [TD3, SAC]: env_id = "Pendulum-v0"