mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-07-26 19:52:45 +00:00
Pass gamma to n-step replay buffer
This commit is contained in:
parent
461f93609c
commit
9c77a6d9ec
2 changed files with 4 additions and 2 deletions
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Reference in a new issue