mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-07-27 20:02:30 +00:00
Enable n-step buffer for SAC and DQN
This commit is contained in:
parent
f2e129b829
commit
f5678e23eb
3 changed files with 13 additions and 9 deletions
|
|
@ -2,7 +2,7 @@ import io
|
|||
import pathlib
|
||||
import time
|
||||
import warnings
|
||||
from typing import Any, Callable, Dict, List, Mapping, Optional, Tuple, Type, Union
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple, Type, Union
|
||||
|
||||
import gym
|
||||
import numpy as np
|
||||
|
|
@ -76,8 +76,8 @@ class OffPolicyAlgorithm(BaseAlgorithm):
|
|||
env: Union[GymEnv, str],
|
||||
policy_base: Type[BasePolicy],
|
||||
learning_rate: Union[float, Callable],
|
||||
replay_buffer_cls: Type[ReplayBuffer] = None,
|
||||
replay_buffer_kwargs: Optional[Mapping[str, Any]] = None,
|
||||
replay_buffer_class: Optional[Type[ReplayBuffer]] = None,
|
||||
replay_buffer_kwargs: Optional[Dict[str, Any]] = None,
|
||||
buffer_size: int = int(1e6),
|
||||
learning_starts: int = 100,
|
||||
batch_size: int = 256,
|
||||
|
|
@ -148,13 +148,13 @@ class OffPolicyAlgorithm(BaseAlgorithm):
|
|||
self.use_sde_at_warmup = use_sde_at_warmup
|
||||
|
||||
self.replay_buffer = None
|
||||
self.replay_buffer_cls = replay_buffer_cls or ReplayBuffer
|
||||
self.replay_buffer_class = replay_buffer_class or ReplayBuffer
|
||||
self.replay_buffer_kwargs = dict(replay_buffer_kwargs or {})
|
||||
|
||||
def _setup_model(self) -> None:
|
||||
self._setup_lr_schedule()
|
||||
self.set_random_seed(self.seed)
|
||||
self.replay_buffer = self.replay_buffer_cls(
|
||||
self.replay_buffer = self.replay_buffer_class(
|
||||
self.buffer_size,
|
||||
self.observation_space,
|
||||
self.action_space,
|
||||
|
|
|
|||
|
|
@ -70,6 +70,8 @@ class DQN(OffPolicyAlgorithm):
|
|||
gradient_steps: int = 1,
|
||||
n_episodes_rollout: int = -1,
|
||||
optimize_memory_usage: bool = False,
|
||||
replay_buffer_class: Optional[Type[ReplayBuffer]] = None,
|
||||
replay_buffer_kwargs: Optional[Dict[str, Any]] = None,
|
||||
target_update_interval: int = 10000,
|
||||
exploration_fraction: float = 0.1,
|
||||
exploration_initial_eps: float = 1.0,
|
||||
|
|
@ -89,8 +91,8 @@ class DQN(OffPolicyAlgorithm):
|
|||
env,
|
||||
DQNPolicy,
|
||||
learning_rate,
|
||||
ReplayBuffer,
|
||||
None,
|
||||
replay_buffer_class,
|
||||
replay_buffer_kwargs,
|
||||
buffer_size,
|
||||
learning_starts,
|
||||
batch_size,
|
||||
|
|
|
|||
|
|
@ -86,6 +86,8 @@ class SAC(OffPolicyAlgorithm):
|
|||
n_episodes_rollout: int = -1,
|
||||
action_noise: Optional[ActionNoise] = None,
|
||||
optimize_memory_usage: bool = False,
|
||||
replay_buffer_class: Optional[Type[ReplayBuffer]] = None,
|
||||
replay_buffer_kwargs: Optional[Dict[str, Any]] = None,
|
||||
ent_coef: Union[str, float] = "auto",
|
||||
target_update_interval: int = 1,
|
||||
target_entropy: Union[str, float] = "auto",
|
||||
|
|
@ -106,8 +108,8 @@ class SAC(OffPolicyAlgorithm):
|
|||
env,
|
||||
SACPolicy,
|
||||
learning_rate,
|
||||
ReplayBuffer,
|
||||
None,
|
||||
replay_buffer_class,
|
||||
replay_buffer_kwargs,
|
||||
buffer_size,
|
||||
learning_starts,
|
||||
batch_size,
|
||||
|
|
|
|||
Loading…
Reference in a new issue