mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-07-27 20:02:30 +00:00
Add A2C
This commit is contained in:
parent
3bc746c6ee
commit
0ad743c85d
5 changed files with 120 additions and 9 deletions
|
|
@ -8,6 +8,7 @@ PyTorch version of [Stable Baselines](https://github.com/hill-a/stable-baselines
|
|||
|
||||
## Implemented Algorithms
|
||||
|
||||
- A2C
|
||||
- CEM-RL (with TD3)
|
||||
- PPO
|
||||
- SAC
|
||||
|
|
@ -18,11 +19,8 @@ PyTorch version of [Stable Baselines](https://github.com/hill-a/stable-baselines
|
|||
|
||||
TODO:
|
||||
- save/load
|
||||
- predict
|
||||
- flexible mlp
|
||||
- logger
|
||||
- better monitor wrapper?
|
||||
- A2C
|
||||
- better predict
|
||||
- complete logger
|
||||
|
||||
Later:
|
||||
- get_parameters / set_parameters
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import os
|
|||
import pytest
|
||||
import numpy as np
|
||||
|
||||
from torchy_baselines import TD3, CEMRL, PPO, SAC
|
||||
from torchy_baselines import A2C, CEMRL, PPO, SAC, TD3
|
||||
from torchy_baselines.common.noise import NormalActionNoise
|
||||
|
||||
|
||||
|
|
@ -28,9 +28,10 @@ def test_cemrl():
|
|||
os.remove("test_save.pth")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_class", [A2C, PPO])
|
||||
@pytest.mark.parametrize("env_id", ['CartPole-v1', 'Pendulum-v0'])
|
||||
def test_ppo(env_id):
|
||||
model = PPO('MlpPolicy', env_id, policy_kwargs=dict(net_arch=[16]), verbose=1, create_eval_env=True)
|
||||
def test_onpolicy(model_class, env_id):
|
||||
model = model_class('MlpPolicy', env_id, policy_kwargs=dict(net_arch=[16]), verbose=1, create_eval_env=True)
|
||||
model.learn(total_timesteps=1000, eval_freq=500)
|
||||
# model.save("test_save")
|
||||
# model.load("test_save")
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
from torchy_baselines.a2c import A2C
|
||||
from torchy_baselines.cem_rl import CEMRL
|
||||
from torchy_baselines.ppo import PPO
|
||||
from torchy_baselines.sac import SAC
|
||||
from torchy_baselines.td3 import TD3
|
||||
|
||||
__version__ = "0.0.4"
|
||||
__version__ = "0.0.5a"
|
||||
|
|
|
|||
2
torchy_baselines/a2c/__init__.py
Normal file
2
torchy_baselines/a2c/__init__.py
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
from torchy_baselines.a2c.a2c import A2C
|
||||
from torchy_baselines.ppo.policies import MlpPolicy
|
||||
109
torchy_baselines/a2c/a2c.py
Normal file
109
torchy_baselines/a2c/a2c.py
Normal file
|
|
@ -0,0 +1,109 @@
|
|||
from gym import spaces
|
||||
import torch as th
|
||||
import torch.nn.functional as F
|
||||
|
||||
from torchy_baselines.common.utils import explained_variance
|
||||
from torchy_baselines.ppo.ppo import PPO
|
||||
from torchy_baselines.ppo.policies import PPOPolicy
|
||||
|
||||
|
||||
class A2C(PPO):
|
||||
"""
|
||||
Advantage Actor Critic (A2C)
|
||||
|
||||
Paper: https://arxiv.org/abs/1602.01783
|
||||
Code: This implementation borrows code from https://github.com/ikostrikov/pytorch-a2c-ppo-acktr-gail and
|
||||
and Stable Baselines (https://github.com/hill-a/stable-baselines)
|
||||
|
||||
Introduction to A2C: https://hackernoon.com/intuitive-rl-intro-to-advantage-actor-critic-a2c-4ff545978752
|
||||
|
||||
:param policy: (PPOPolicy or str) The policy model to use (MlpPolicy, CnnPolicy, ...)
|
||||
:param env: (Gym environment or str) The environment to learn from (if registered in Gym, can be str)
|
||||
:param learning_rate: (float or callable) The learning rate, it can be a function
|
||||
:param n_steps: (int) The number of steps to run for each environment per update
|
||||
(i.e. batch size is n_steps * n_env where n_env is number of environment copies running in parallel)
|
||||
:param batch_size: (int) Minibatch size
|
||||
:param n_epochs: (int) Number of epoch when optimizing the surrogate loss
|
||||
:param gamma: (float) Discount factor
|
||||
:param gae_lambda: (float) Factor for trade-off of bias vs variance for Generalized Advantage Estimator
|
||||
:param ent_coef: (float) Entropy coefficient for the loss calculation
|
||||
:param vf_coef: (float) Value function coefficient for the loss calculation
|
||||
:param max_grad_norm: (float) The maximum value for the gradient clipping
|
||||
:param tensorboard_log: (str) the log location for tensorboard (if None, no logging)
|
||||
:param create_eval_env: (bool) Whether to create a second environment that will be
|
||||
used for evaluating the agent periodically. (Only available when passing string for the environment)
|
||||
:param policy_kwargs: (dict) additional arguments to be passed to the policy on creation
|
||||
:param verbose: (int) the verbosity level: 0 none, 1 training information, 2 tensorflow debug
|
||||
:param seed: (int) Seed for the pseudo random generators
|
||||
:param device: (str or th.device) Device (cpu, cuda, ...) on which the code should be run.
|
||||
Setting it to auto, the code will be run on the GPU if possible.
|
||||
:param _init_setup_model: (bool) Whether or not to build the network at the creation of the instance
|
||||
"""
|
||||
|
||||
def __init__(self, policy, env, learning_rate=3e-4,
|
||||
n_steps=2048, batch_size=64, n_epochs=1,
|
||||
gamma=0.99, gae_lambda=0.95,
|
||||
ent_coef=0.0, vf_coef=0.5, max_grad_norm=0.5,
|
||||
tensorboard_log=None, create_eval_env=False,
|
||||
policy_kwargs=None, verbose=0, seed=0, device='auto',
|
||||
_init_setup_model=True):
|
||||
|
||||
super(A2C, self).__init__(policy, env, learning_rate=learning_rate,
|
||||
n_steps=n_steps, batch_size=batch_size, n_epochs=n_epochs,
|
||||
gamma=gamma, gae_lambda=gae_lambda, ent_coef=ent_coef,
|
||||
vf_coef=vf_coef, max_grad_norm=max_grad_norm,
|
||||
tensorboard_log=tensorboard_log, policy_kwargs=policy_kwargs,
|
||||
verbose=verbose, device=device, create_eval_env=create_eval_env,
|
||||
seed=seed, _init_setup_model=False)
|
||||
|
||||
self.batch_size = n_steps
|
||||
|
||||
if _init_setup_model:
|
||||
self._setup_model()
|
||||
|
||||
def train(self, gradient_steps, batch_size=64):
|
||||
|
||||
for gradient_step in range(gradient_steps):
|
||||
# approx_kl_divs = []
|
||||
# Sample replay buffer
|
||||
for replay_data in self.rollout_buffer.get(batch_size):
|
||||
# Unpack
|
||||
obs, action, _, _, advantage, return_batch = replay_data
|
||||
|
||||
if isinstance(self.action_space, spaces.Discrete):
|
||||
# Convert discrete action for float to long
|
||||
action = action.long().flatten()
|
||||
|
||||
values, log_prob, entropy = self.policy.get_policy_stats(obs, action)
|
||||
values = values.flatten()
|
||||
# Normalize advantage
|
||||
# TODO: check without
|
||||
advantage = (advantage - advantage.mean()) / (advantage.std() + 1e-8)
|
||||
|
||||
policy_loss = -(advantage * log_prob).mean()
|
||||
|
||||
# Value loss using the TD(gae_lambda) target
|
||||
value_loss = F.mse_loss(return_batch, values)
|
||||
|
||||
# Entropy loss favor exploration
|
||||
entropy_loss = th.mean(entropy)
|
||||
|
||||
loss = policy_loss + self.ent_coef * entropy_loss + self.vf_coef * value_loss
|
||||
|
||||
# Optimization step
|
||||
self.policy.optimizer.zero_grad()
|
||||
loss.backward()
|
||||
# Clip grad norm
|
||||
th.nn.utils.clip_grad_norm_(self.policy.parameters(), self.max_grad_norm)
|
||||
self.policy.optimizer.step()
|
||||
# approx_kl_divs.append(th.mean(old_log_prob - log_prob).detach().cpu().numpy())
|
||||
|
||||
# print(explained_variance(self.rollout_buffer.returns.flatten().cpu().numpy(),
|
||||
# self.rollout_buffer.values.flatten().cpu().numpy()))
|
||||
|
||||
def learn(self, total_timesteps, callback=None, log_interval=100,
|
||||
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="A2C", reset_num_timesteps=True):
|
||||
|
||||
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,
|
||||
tb_log_name=tb_log_name, reset_num_timesteps=reset_num_timesteps)
|
||||
Loading…
Reference in a new issue