mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-09-16 22:20:26 +00:00
Finish typing A2C and PPO
This commit is contained in:
parent
90d1558534
commit
c5e5812894
3 changed files with 39 additions and 16 deletions
|
|
@ -7,11 +7,11 @@ import torch.nn.functional as F
|
|||
from torchy_baselines.common.utils import explained_variance
|
||||
from torchy_baselines.common import logger
|
||||
from torchy_baselines.common.type_aliases import GymEnv
|
||||
from torchy_baselines.common.callbacks import BaseCallback
|
||||
from torchy_baselines.ppo.ppo import PPO
|
||||
from torchy_baselines.ppo.policies import PPOPolicy
|
||||
|
||||
|
||||
|
||||
class A2C(PPO):
|
||||
"""
|
||||
Advantage Actor Critic (A2C)
|
||||
|
|
@ -89,14 +89,14 @@ class A2C(PPO):
|
|||
if _init_setup_model:
|
||||
self._setup_model()
|
||||
|
||||
def _setup_model(self):
|
||||
def _setup_model(self) -> None:
|
||||
super(A2C, self)._setup_model()
|
||||
if self.use_rms_prop:
|
||||
self.policy.optimizer = th.optim.RMSprop(self.policy.parameters(),
|
||||
lr=self.learning_rate(1), alpha=0.99,
|
||||
eps=self.rms_prop_eps, weight_decay=0)
|
||||
|
||||
def train(self, gradient_steps: int, batch_size=None):
|
||||
def train(self, gradient_steps: int, batch_size: Optional[int] = None) -> None:
|
||||
# Update optimizer learning rate
|
||||
self._update_learning_rate(self.policy.optimizer)
|
||||
# A2C with gradient_steps > 1 does not make sense
|
||||
|
|
@ -153,9 +153,16 @@ class A2C(PPO):
|
|||
if hasattr(self.policy, 'log_std'):
|
||||
logger.logkv("std", th.exp(self.policy.log_std).mean().item())
|
||||
|
||||
def learn(self, total_timesteps, callback=None, log_interval=100,
|
||||
eval_env=None, eval_freq=-1, n_eval_episodes=5,
|
||||
tb_log_name="A2C", eval_log_path=None, reset_num_timesteps=True):
|
||||
def learn(self,
|
||||
total_timesteps: int,
|
||||
callback: Optional[BaseCallback] = None,
|
||||
log_interval: int = 100,
|
||||
eval_env: Optional[GymEnv] = None,
|
||||
eval_freq: int = -1,
|
||||
n_eval_episodes: int = 5,
|
||||
tb_log_name: str = "A2C",
|
||||
eval_log_path: Optional[str] = None,
|
||||
reset_num_timesteps: bool = True) -> 'A2C':
|
||||
|
||||
return super(A2C, self).learn(total_timesteps=total_timesteps, callback=callback, log_interval=log_interval,
|
||||
eval_env=eval_env, eval_freq=eval_freq, n_eval_episodes=n_eval_episodes,
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
from typing import Type, Tuple, Optional, List
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
|
|
@ -21,9 +23,16 @@ class CEM(object):
|
|||
:param antithetic: (bool) Use a finite difference like method for sampling
|
||||
(mu + epsilon, mu - epsilon)
|
||||
"""
|
||||
def __init__(self, num_params, mu_init=None, sigma_init=1e-3,
|
||||
pop_size=256, damping_init=1e-3, damping_final=1e-5,
|
||||
parents=None, elitism=False, antithetic=False):
|
||||
def __init__(self,
|
||||
num_params: int,
|
||||
mu_init: Optional[np.ndarray] = None,
|
||||
sigma_init: float = 1e-3,
|
||||
pop_size: int = 256,
|
||||
damping_init: float = 1e-3,
|
||||
damping_final: float = 1e-5,
|
||||
parents: Optional[int] = None,
|
||||
elitism: bool = False,
|
||||
antithetic: bool = False):
|
||||
super(CEM, self).__init__()
|
||||
|
||||
self.num_params = num_params
|
||||
|
|
@ -66,7 +75,7 @@ class CEM(object):
|
|||
for i in range(1, self.parents + 1)])
|
||||
self.weights /= self.weights.sum()
|
||||
|
||||
def ask(self, pop_size):
|
||||
def ask(self, pop_size: int) -> List[np.ndarray]:
|
||||
"""
|
||||
Returns a list of candidates parameters
|
||||
|
||||
|
|
@ -87,7 +96,7 @@ class CEM(object):
|
|||
|
||||
return individuals
|
||||
|
||||
def tell(self, solutions, scores):
|
||||
def tell(self, solutions: List[np.ndarray], scores: List[float]) -> None:
|
||||
"""
|
||||
Updates the distribution
|
||||
|
||||
|
|
@ -114,7 +123,7 @@ class CEM(object):
|
|||
self.elite = solutions[idx_sorted[0]]
|
||||
self.elite_score = scores[idx_sorted[0]]
|
||||
|
||||
def get_distrib_params(self):
|
||||
def get_distrib_params(self) -> Tuple[np.ndarray, np.ndarray]:
|
||||
"""
|
||||
Returns the parameters of the distribution:
|
||||
the mean and standard deviation.
|
||||
|
|
|
|||
|
|
@ -121,7 +121,7 @@ class PPO(BaseRLModel):
|
|||
if _init_setup_model:
|
||||
self._setup_model()
|
||||
|
||||
def _setup_model(self):
|
||||
def _setup_model(self) -> None:
|
||||
self._setup_learning_rate()
|
||||
# TODO: preprocessing: one hot vector for obs discrete
|
||||
state_dim = self.observation_space.shape[0]
|
||||
|
|
@ -284,9 +284,16 @@ class PPO(BaseRLModel):
|
|||
if hasattr(self.policy, 'log_std'):
|
||||
logger.logkv("std", th.exp(self.policy.log_std).mean().item())
|
||||
|
||||
def learn(self, total_timesteps, callback=None, log_interval=1,
|
||||
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="PPO",
|
||||
eval_log_path=None, reset_num_timesteps=True):
|
||||
def learn(self,
|
||||
total_timesteps: int,
|
||||
callback: Optional[BaseCallback] = None,
|
||||
log_interval: int = 1,
|
||||
eval_env: Optional[GymEnv] = None,
|
||||
eval_freq: int = -1,
|
||||
n_eval_episodes: int = 5,
|
||||
tb_log_name: str = "PPO",
|
||||
eval_log_path: Optional[str] = None,
|
||||
reset_num_timesteps: bool = True) -> 'PPO':
|
||||
|
||||
episode_num, obs, callback = self._setup_learn(eval_env, callback, eval_freq,
|
||||
n_eval_episodes, eval_log_path, reset_num_timesteps)
|
||||
|
|
|
|||
Loading…
Reference in a new issue