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.utils import explained_variance
|
||||||
from torchy_baselines.common import logger
|
from torchy_baselines.common import logger
|
||||||
from torchy_baselines.common.type_aliases import GymEnv
|
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.ppo import PPO
|
||||||
from torchy_baselines.ppo.policies import PPOPolicy
|
from torchy_baselines.ppo.policies import PPOPolicy
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class A2C(PPO):
|
class A2C(PPO):
|
||||||
"""
|
"""
|
||||||
Advantage Actor Critic (A2C)
|
Advantage Actor Critic (A2C)
|
||||||
|
|
@ -89,14 +89,14 @@ class A2C(PPO):
|
||||||
if _init_setup_model:
|
if _init_setup_model:
|
||||||
self._setup_model()
|
self._setup_model()
|
||||||
|
|
||||||
def _setup_model(self):
|
def _setup_model(self) -> None:
|
||||||
super(A2C, self)._setup_model()
|
super(A2C, self)._setup_model()
|
||||||
if self.use_rms_prop:
|
if self.use_rms_prop:
|
||||||
self.policy.optimizer = th.optim.RMSprop(self.policy.parameters(),
|
self.policy.optimizer = th.optim.RMSprop(self.policy.parameters(),
|
||||||
lr=self.learning_rate(1), alpha=0.99,
|
lr=self.learning_rate(1), alpha=0.99,
|
||||||
eps=self.rms_prop_eps, weight_decay=0)
|
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
|
# Update optimizer learning rate
|
||||||
self._update_learning_rate(self.policy.optimizer)
|
self._update_learning_rate(self.policy.optimizer)
|
||||||
# A2C with gradient_steps > 1 does not make sense
|
# A2C with gradient_steps > 1 does not make sense
|
||||||
|
|
@ -153,9 +153,16 @@ class A2C(PPO):
|
||||||
if hasattr(self.policy, 'log_std'):
|
if hasattr(self.policy, 'log_std'):
|
||||||
logger.logkv("std", th.exp(self.policy.log_std).mean().item())
|
logger.logkv("std", th.exp(self.policy.log_std).mean().item())
|
||||||
|
|
||||||
def learn(self, total_timesteps, callback=None, log_interval=100,
|
def learn(self,
|
||||||
eval_env=None, eval_freq=-1, n_eval_episodes=5,
|
total_timesteps: int,
|
||||||
tb_log_name="A2C", eval_log_path=None, reset_num_timesteps=True):
|
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,
|
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,
|
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
|
import numpy as np
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -21,9 +23,16 @@ class CEM(object):
|
||||||
:param antithetic: (bool) Use a finite difference like method for sampling
|
:param antithetic: (bool) Use a finite difference like method for sampling
|
||||||
(mu + epsilon, mu - epsilon)
|
(mu + epsilon, mu - epsilon)
|
||||||
"""
|
"""
|
||||||
def __init__(self, num_params, mu_init=None, sigma_init=1e-3,
|
def __init__(self,
|
||||||
pop_size=256, damping_init=1e-3, damping_final=1e-5,
|
num_params: int,
|
||||||
parents=None, elitism=False, antithetic=False):
|
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__()
|
super(CEM, self).__init__()
|
||||||
|
|
||||||
self.num_params = num_params
|
self.num_params = num_params
|
||||||
|
|
@ -66,7 +75,7 @@ class CEM(object):
|
||||||
for i in range(1, self.parents + 1)])
|
for i in range(1, self.parents + 1)])
|
||||||
self.weights /= self.weights.sum()
|
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
|
Returns a list of candidates parameters
|
||||||
|
|
||||||
|
|
@ -87,7 +96,7 @@ class CEM(object):
|
||||||
|
|
||||||
return individuals
|
return individuals
|
||||||
|
|
||||||
def tell(self, solutions, scores):
|
def tell(self, solutions: List[np.ndarray], scores: List[float]) -> None:
|
||||||
"""
|
"""
|
||||||
Updates the distribution
|
Updates the distribution
|
||||||
|
|
||||||
|
|
@ -114,7 +123,7 @@ class CEM(object):
|
||||||
self.elite = solutions[idx_sorted[0]]
|
self.elite = solutions[idx_sorted[0]]
|
||||||
self.elite_score = scores[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:
|
Returns the parameters of the distribution:
|
||||||
the mean and standard deviation.
|
the mean and standard deviation.
|
||||||
|
|
|
||||||
|
|
@ -121,7 +121,7 @@ class PPO(BaseRLModel):
|
||||||
if _init_setup_model:
|
if _init_setup_model:
|
||||||
self._setup_model()
|
self._setup_model()
|
||||||
|
|
||||||
def _setup_model(self):
|
def _setup_model(self) -> None:
|
||||||
self._setup_learning_rate()
|
self._setup_learning_rate()
|
||||||
# TODO: preprocessing: one hot vector for obs discrete
|
# TODO: preprocessing: one hot vector for obs discrete
|
||||||
state_dim = self.observation_space.shape[0]
|
state_dim = self.observation_space.shape[0]
|
||||||
|
|
@ -284,9 +284,16 @@ class PPO(BaseRLModel):
|
||||||
if hasattr(self.policy, 'log_std'):
|
if hasattr(self.policy, 'log_std'):
|
||||||
logger.logkv("std", th.exp(self.policy.log_std).mean().item())
|
logger.logkv("std", th.exp(self.policy.log_std).mean().item())
|
||||||
|
|
||||||
def learn(self, total_timesteps, callback=None, log_interval=1,
|
def learn(self,
|
||||||
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="PPO",
|
total_timesteps: int,
|
||||||
eval_log_path=None, reset_num_timesteps=True):
|
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,
|
episode_num, obs, callback = self._setup_learn(eval_env, callback, eval_freq,
|
||||||
n_eval_episodes, eval_log_path, reset_num_timesteps)
|
n_eval_episodes, eval_log_path, reset_num_timesteps)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue