stable-baselines3/torchy_baselines/sac/sac.py

314 lines
16 KiB
Python
Raw Normal View History

from typing import List, Tuple
2019-09-24 12:15:12 +00:00
import torch as th
import torch.nn.functional as F
import numpy as np
from torchy_baselines.common.base_class import BaseRLModel
from torchy_baselines.common.buffers import ReplayBuffer
from torchy_baselines.sac.policies import SACPolicy
2019-11-26 16:44:06 +00:00
from torchy_baselines.common import logger
2019-09-24 12:15:12 +00:00
class SAC(BaseRLModel):
"""
2019-09-24 13:30:58 +00:00
Soft Actor-Critic (SAC)
2019-09-24 12:15:12 +00:00
Off-Policy Maximum Entropy Deep Reinforcement Learning with a Stochastic Actor,
2019-09-24 13:30:58 +00:00
This implementation borrows code from original implementation (https://github.com/haarnoja/sac)
from OpenAI Spinning Up (https://github.com/openai/spinningup), from the softlearning repo
2019-09-24 12:15:12 +00:00
(https://github.com/rail-berkeley/softlearning/)
2019-09-24 13:30:58 +00:00
and from Stable Baselines (https://github.com/hill-a/stable-baselines)
Paper: https://arxiv.org/abs/1801.01290
Introduction to SAC: https://spinningup.openai.com/en/latest/algorithms/sac.html
2019-09-24 12:15:12 +00:00
Note: we use double q target and not value target as discussed
in https://github.com/hill-a/stable-baselines/issues/270
2019-09-24 13:30:58 +00:00
:param policy: (SACPolicy 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) learning rate for adam optimizer,
the same learning rate will be used for all networks (Q-Values, Actor and Value function)
it can be a function of the current progress (from 1 to 0)
:param buffer_size: (int) size of the replay buffer
:param batch_size: (int) Minibatch size for each gradient update
:param tau: (float) the soft update coefficient ("polyak update", between 0 and 1)
:param ent_coef: (str or float) Entropy regularization coefficient. (Equivalent to
inverse of reward scale in the original SAC paper.) Controlling exploration/exploitation trade-off.
Set it to 'auto' to learn it automatically (and 'auto_0.1' for using 0.1 as initial value)
:param learning_starts: (int) how many steps of the model to collect transitions for before learning starts
:param target_update_interval: (int) update the target network every `target_network_update_freq` steps.
2019-09-25 11:20:06 +00:00
:param train_freq: (int) Update the model every `train_freq` steps.
2019-09-24 13:30:58 +00:00
:param gradient_steps: (int) How many gradient update after each step
2020-01-20 15:19:35 +00:00
:param n_episodes_rollout: (int) Update the model every `n_episodes_rollout` episodes.
Note that this cannot be used at the same time as `train_freq`
2020-01-22 16:17:12 +00:00
:param target_entropy: (str or float) target entropy when learning `ent_coef` (`ent_coef = 'auto'`)
2019-09-24 13:30:58 +00:00
:param action_noise: (ActionNoise) the action noise type (None by default), this can help
2019-10-07 14:26:03 +00:00
for hard exploration problem. Cf common.noise for the different action noise type.
2019-09-24 13:30:58 +00:00
:param gamma: (float) the discount factor
2019-11-26 14:26:12 +00:00
:param use_sde: (bool) Whether to use State Dependent Exploration (SDE)
instead of action noise exploration (default: False)
2019-12-17 10:47:21 +00:00
:param sde_sample_freq: (int) Sample a new noise matrix every n steps when using SDE
Default: -1 (only sample at the beginning of the rollout)
2019-09-24 13:30:58 +00:00
: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
2019-09-26 09:46:40 +00:00
: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.
2019-09-24 13:30:58 +00:00
:param _init_setup_model: (bool) Whether or not to build the network at the creation of the instance
"""
2019-11-28 14:38:04 +00:00
2019-09-24 13:30:58 +00:00
def __init__(self, policy, env, learning_rate=3e-4, buffer_size=int(1e6),
2019-11-14 13:35:47 +00:00
learning_starts=100, batch_size=256,
2019-09-24 13:30:58 +00:00
tau=0.005, ent_coef='auto', target_update_interval=1,
2019-09-25 15:07:54 +00:00
train_freq=1, gradient_steps=1, n_episodes_rollout=-1,
2019-12-18 15:56:51 +00:00
target_entropy='auto', action_noise=None,
2019-12-17 10:47:21 +00:00
gamma=0.99, use_sde=False, sde_sample_freq=-1,
tensorboard_log=None, create_eval_env=False,
2019-09-24 14:59:47 +00:00
policy_kwargs=None, verbose=0, seed=0, device='auto',
2019-09-24 12:15:12 +00:00
_init_setup_model=True):
super(SAC, self).__init__(policy, env, SACPolicy, policy_kwargs, verbose, device,
2020-01-20 10:17:55 +00:00
create_eval_env=create_eval_env, seed=seed,
use_sde=use_sde, sde_sample_freq=sde_sample_freq)
2019-09-24 12:15:12 +00:00
self.learning_rate = learning_rate
self.target_entropy = target_entropy
self.log_ent_coef = None
2019-09-25 11:20:06 +00:00
self.target_update_interval = target_update_interval
2019-09-24 13:30:58 +00:00
self.buffer_size = buffer_size
# In the original paper, same learning rate is used for all networks
self.learning_rate = learning_rate
self.learning_starts = learning_starts
self.batch_size = batch_size
self.tau = tau
# Entropy coefficient / Entropy temperature
# Inverse of the reward scale
self.ent_coef = ent_coef
self.target_update_interval = target_update_interval
2019-09-25 11:20:06 +00:00
self.train_freq = train_freq
self.gradient_steps = gradient_steps
self.n_episodes_rollout = n_episodes_rollout
2019-10-07 14:26:03 +00:00
self.action_noise = action_noise
2019-09-24 12:15:12 +00:00
self.gamma = gamma
self.ent_coef_optimizer = None
2019-09-24 12:15:12 +00:00
if _init_setup_model:
self._setup_model()
def _setup_model(self):
2019-10-28 15:47:13 +00:00
self._setup_learning_rate()
2019-09-24 12:53:03 +00:00
obs_dim, action_dim = self.observation_space.shape[0], self.action_space.shape[0]
2019-10-10 11:47:13 +00:00
if self.seed is not None:
self.set_random_seed(self.seed)
2019-09-24 12:15:12 +00:00
# Target entropy is used when learning the entropy coefficient
if self.target_entropy == 'auto':
# automatically set target entropy if needed
self.target_entropy = -np.prod(self.env.action_space.shape).astype(np.float32)
else:
# Force conversion
# this will also throw an error for unexpected string
self.target_entropy = float(self.target_entropy)
# The entropy coefficient or entropy can be learned automatically
# see Automating Entropy Adjustment for Maximum Entropy RL section
# of https://arxiv.org/abs/1812.05905
if isinstance(self.ent_coef, str) and self.ent_coef.startswith('auto'):
# Default initial value of ent_coef when learned
init_value = 1.0
if '_' in self.ent_coef:
init_value = float(self.ent_coef.split('_')[1])
assert init_value > 0., "The initial value of ent_coef must be greater than 0"
# Note: we optimize the log of the entropy coeff which is slightly different from the paper
# as discussed in https://github.com/rail-berkeley/softlearning/issues/37
self.log_ent_coef = th.log(th.ones(1, device=self.device) * init_value).requires_grad_(True)
2019-10-28 15:47:13 +00:00
self.ent_coef_optimizer = th.optim.Adam([self.log_ent_coef], lr=self.learning_rate(1))
2019-09-24 12:15:12 +00:00
else:
# Force conversion to float
# this will throw an error if a malformed string (different from 'auto')
# is passed
self.ent_coef_tensor = th.tensor(float(self.ent_coef)).to(self.device)
2019-09-24 12:15:12 +00:00
2019-09-24 12:53:03 +00:00
self.replay_buffer = ReplayBuffer(self.buffer_size, obs_dim, action_dim, self.device)
self.policy = self.policy_class(self.observation_space, self.action_space,
2020-01-20 10:17:55 +00:00
self.learning_rate, use_sde=self.use_sde,
device=self.device, **self.policy_kwargs)
2019-09-24 12:15:12 +00:00
self.policy = self.policy.to(self.device)
self._create_aliases()
def _create_aliases(self):
self.actor = self.policy.actor
self.critic = self.policy.critic
self.critic_target = self.policy.critic_target
def select_action(self, observation):
# Normally not needed
observation = np.array(observation)
with th.no_grad():
observation = th.FloatTensor(observation.reshape(1, -1)).to(self.device)
return self.actor(observation).cpu().data.numpy()
def predict(self, observation, state=None, mask=None, deterministic=True):
"""
Get the model's action from an observation
:param observation: (np.ndarray) the input observation
:param state: (np.ndarray) The last states (can be None, used in recurrent policies)
:param mask: (np.ndarray) The last masks (can be None, used in recurrent policies)
:param deterministic: (bool) Whether or not to return deterministic actions.
:return: (np.ndarray, np.ndarray) the model's action and the next state (used in recurrent policies)
"""
2019-10-07 14:26:03 +00:00
return self.unscale_action(self.select_action(observation))
2019-09-24 12:15:12 +00:00
2020-01-22 16:17:12 +00:00
def train(self, gradient_steps: int, batch_size: int = 64):
2019-10-28 15:47:13 +00:00
# Update optimizers learning rate
optimizers = [self.actor.optimizer, self.critic.optimizer]
if self.ent_coef_optimizer is not None:
optimizers += [self.ent_coef_optimizer]
self._update_learning_rate(optimizers)
2020-01-22 16:17:12 +00:00
ent_coef_loss, ent_coef = th.zeros(1), th.zeros(1)
actor_loss, critic_loss = th.zeros(1), th.zeros(1)
2019-09-25 11:20:06 +00:00
for gradient_step in range(gradient_steps):
2019-09-24 12:15:12 +00:00
# Sample replay buffer
2019-11-14 13:35:00 +00:00
replay_data = self.replay_buffer.sample(batch_size, env=self._vec_normalize_env)
2019-09-24 12:15:12 +00:00
2019-09-24 12:53:03 +00:00
obs, action_batch, next_obs, done, reward = replay_data
2019-09-24 12:15:12 +00:00
# Two options: retain_graph=True in the actor_loss.backward()
# or sample again the noise matrix
# otherwise the intermediate step `std = th.exp(log_std)`
# is lost and we cannot backpropagate through again
2019-12-02 13:06:30 +00:00
# anyway, we need to sample because `log_std` may have changed between two gradient steps
if self.use_sde:
self.actor.reset_noise(batch_size=batch_size)
# self.actor.reset_noise()
2019-09-24 12:15:12 +00:00
# Action by the current actor for the sampled state
2019-09-24 12:53:03 +00:00
action_pi, log_prob = self.actor.action_log_prob(obs)
2019-09-24 12:15:12 +00:00
log_prob = log_prob.reshape(-1, 1)
ent_coef_loss = None
if self.ent_coef_optimizer is not None:
2019-09-25 11:30:08 +00:00
# Important: detach the variable from the graph
# so we don't change it with other losses
# see https://github.com/rail-berkeley/softlearning/issues/60
ent_coef = th.exp(self.log_ent_coef.detach())
2019-09-24 12:15:12 +00:00
ent_coef_loss = -(self.log_ent_coef * (log_prob + self.target_entropy).detach()).mean()
2019-09-25 11:30:08 +00:00
else:
ent_coef = self.ent_coef_tensor
2019-09-24 12:15:12 +00:00
# Optimize entropy coefficient, also called
# entropy temperature or alpha in the paper
if ent_coef_loss is not None:
self.ent_coef_optimizer.zero_grad()
ent_coef_loss.backward()
self.ent_coef_optimizer.step()
2019-09-25 11:30:08 +00:00
with th.no_grad():
2019-12-02 13:06:30 +00:00
# if self.use_sde:
# self.actor.reset_noise(batch_size=batch_size)
# Select action according to policy
next_action, next_log_prob = self.actor.action_log_prob(next_obs)
2019-09-25 11:30:08 +00:00
# Compute the target Q value
target_q1, target_q2 = self.critic_target(next_obs, next_action)
target_q = th.min(target_q1, target_q2)
target_q = reward + (1 - done) * self.gamma * target_q
# td error + entropy term
q_backup = target_q - ent_coef * next_log_prob.reshape(-1, 1)
2019-09-24 12:15:12 +00:00
# Get current Q estimates
# using action from the replay buffer
2019-09-24 12:53:03 +00:00
current_q1, current_q2 = self.critic(obs, action_batch)
2019-09-24 12:15:12 +00:00
# Compute critic loss
critic_loss = 0.5 * (F.mse_loss(current_q1, q_backup) + F.mse_loss(current_q2, q_backup))
# Optimize the critic
self.critic.optimizer.zero_grad()
critic_loss.backward()
self.critic.optimizer.step()
# Compute actor loss
2019-09-25 11:30:08 +00:00
# Alternative: actor_loss = th.mean(log_prob - qf1_pi)
qf1_pi, qf2_pi = self.critic.forward(obs, action_pi)
min_qf_pi = th.min(qf1_pi, qf2_pi)
actor_loss = (ent_coef * log_prob - min_qf_pi).mean()
2019-09-24 12:15:12 +00:00
# Optimize the actor
self.actor.optimizer.zero_grad()
2019-12-02 13:06:30 +00:00
actor_loss.backward()
2019-09-24 12:15:12 +00:00
self.actor.optimizer.step()
# Update target networks
2019-09-25 11:20:06 +00:00
if gradient_step % self.target_update_interval == 0:
for param, target_param in zip(self.critic.parameters(), self.critic_target.parameters()):
target_param.data.copy_(self.tau * param.data + (1 - self.tau) * target_param.data)
2019-09-24 12:15:12 +00:00
2019-11-26 16:44:06 +00:00
# TODO: average
logger.logkv("ent_coef", ent_coef.item())
logger.logkv("actor_loss", actor_loss.item())
logger.logkv("critic_loss", critic_loss.item())
if ent_coef_loss is not None:
logger.logkv("ent_coef_loss", ent_coef_loss.item())
2019-11-26 16:44:06 +00:00
2019-10-10 11:47:13 +00:00
def learn(self, total_timesteps, callback=None, log_interval=4,
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="SAC",
2020-01-27 14:53:27 +00:00
eval_log_path=None, reset_num_timesteps=True):
2019-09-24 12:15:12 +00:00
2020-01-31 12:16:28 +00:00
episode_num, obs, callback = self._setup_learn(eval_env, callback, eval_freq,
n_eval_episodes, eval_log_path, reset_num_timesteps)
2020-01-27 13:32:31 +00:00
callback.on_training_start(locals(), globals())
2019-09-24 12:15:12 +00:00
2020-01-27 13:32:31 +00:00
while self.num_timesteps < total_timesteps:
2019-09-25 11:20:06 +00:00
rollout = self.collect_rollouts(self.env, n_episodes=self.n_episodes_rollout,
2019-10-07 14:26:03 +00:00
n_steps=self.train_freq, action_noise=self.action_noise,
2020-01-27 13:32:31 +00:00
deterministic=False, callback=callback,
2019-09-25 11:20:06 +00:00
learning_starts=self.learning_starts,
replay_buffer=self.replay_buffer,
2019-10-10 11:47:13 +00:00
obs=obs, episode_num=episode_num,
log_interval=log_interval)
2019-09-25 11:20:06 +00:00
# Unpack
2020-01-27 13:32:31 +00:00
episode_reward, episode_timesteps, n_episodes, obs, continue_training = rollout
if continue_training is False:
break
2019-09-25 11:20:06 +00:00
episode_num += n_episodes
2019-10-28 15:47:13 +00:00
self._update_current_progress(self.num_timesteps, total_timesteps)
2019-09-24 12:15:12 +00:00
2019-10-01 19:56:37 +00:00
if self.num_timesteps > 0 and self.num_timesteps > self.learning_starts:
2019-09-25 11:20:06 +00:00
gradient_steps = self.gradient_steps if self.gradient_steps > 0 else episode_timesteps
self.train(gradient_steps, batch_size=self.batch_size)
2019-09-24 12:15:12 +00:00
2020-01-27 13:32:31 +00:00
callback.on_training_end()
2019-09-24 12:15:12 +00:00
return self
def excluded_save_params(self) -> List[str]:
2019-11-21 15:46:53 +00:00
"""
Returns the names of the parameters that should be excluded by default
when saving the model.
2019-11-21 15:46:53 +00:00
:return: (List[str]) List of parameters that should be excluded from save
2019-11-21 15:46:53 +00:00
"""
# Exclude aliases
return super(SAC, self).excluded_save_params() + ["actor", "critic", "critic_target"]
2019-11-21 15:46:53 +00:00
def get_torch_variables(self) -> Tuple[List[str], List[str]]:
2019-11-21 15:46:53 +00:00
"""
cf base class
2019-11-21 15:46:53 +00:00
"""
state_dicts = ["policy", "actor.optimizer", "critic.optimizer"]
saved_tensors = ['log_ent_coef']
if self.ent_coef_optimizer is not None:
state_dicts.append('ent_coef_optimizer')
else:
saved_tensors.append('ent_coef_tensor')
return state_dicts, saved_tensors