Finish typing A2C and PPO

This commit is contained in:
Antonin Raffin 2020-03-11 13:01:42 +01:00
parent 90d1558534
commit c5e5812894
3 changed files with 39 additions and 16 deletions

View file

@ -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,

View file

@ -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.

View file

@ -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)