mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-09-16 22:20:26 +00:00
Add docstrings
This commit is contained in:
parent
f4fe1362f0
commit
32648d9029
6 changed files with 130 additions and 59 deletions
|
|
@ -5,7 +5,7 @@ from torchy_baselines import TD3, CEMRL, PPO, SAC
|
||||||
|
|
||||||
def test_td3():
|
def test_td3():
|
||||||
model = TD3('MlpPolicy', 'Pendulum-v0', policy_kwargs=dict(net_arch=[64, 64]),
|
model = TD3('MlpPolicy', 'Pendulum-v0', policy_kwargs=dict(net_arch=[64, 64]),
|
||||||
start_timesteps=100, verbose=1, create_eval_env=True)
|
learning_starts=100, verbose=1, create_eval_env=True)
|
||||||
model.learn(total_timesteps=1000, eval_freq=500)
|
model.learn(total_timesteps=1000, eval_freq=500)
|
||||||
model.save("test_save")
|
model.save("test_save")
|
||||||
model.load("test_save")
|
model.load("test_save")
|
||||||
|
|
@ -14,7 +14,7 @@ def test_td3():
|
||||||
|
|
||||||
def test_cemrl():
|
def test_cemrl():
|
||||||
model = CEMRL('MlpPolicy', 'Pendulum-v0', policy_kwargs=dict(net_arch=[16]), pop_size=2, n_grad=1,
|
model = CEMRL('MlpPolicy', 'Pendulum-v0', policy_kwargs=dict(net_arch=[16]), pop_size=2, n_grad=1,
|
||||||
start_timesteps=100, verbose=1, create_eval_env=True)
|
learning_starts=100, verbose=1, create_eval_env=True)
|
||||||
model.learn(total_timesteps=1000, eval_freq=500)
|
model.learn(total_timesteps=1000, eval_freq=500)
|
||||||
model.save("test_save")
|
model.save("test_save")
|
||||||
model.load("test_save")
|
model.load("test_save")
|
||||||
|
|
@ -30,5 +30,5 @@ def test_ppo():
|
||||||
|
|
||||||
def test_sac():
|
def test_sac():
|
||||||
model = SAC('MlpPolicy', 'Pendulum-v0', policy_kwargs=dict(net_arch=[64, 64]),
|
model = SAC('MlpPolicy', 'Pendulum-v0', policy_kwargs=dict(net_arch=[64, 64]),
|
||||||
start_timesteps=100, verbose=1, create_eval_env=True, ent_coef='auto')
|
learning_starts=100, verbose=1, create_eval_env=True, ent_coef='auto')
|
||||||
model.learn(total_timesteps=1000, eval_freq=500)
|
model.learn(total_timesteps=1000, eval_freq=500)
|
||||||
|
|
|
||||||
|
|
@ -17,16 +17,17 @@ class CEMRL(TD3):
|
||||||
|
|
||||||
def __init__(self, policy, env, policy_kwargs=None, verbose=0,
|
def __init__(self, policy, env, policy_kwargs=None, verbose=0,
|
||||||
sigma_init=1e-3, pop_size=10, damp=1e-3, damp_limit=1e-5,
|
sigma_init=1e-3, pop_size=10, damp=1e-3, damp_limit=1e-5,
|
||||||
elitism=False, n_grad=5, policy_freq=2, batch_size=100,
|
elitism=False, n_grad=5, policy_delay=2, batch_size=100,
|
||||||
buffer_size=int(1e6), learning_rate=1e-3, seed=0, device='auto',
|
buffer_size=int(1e6), learning_rate=1e-3, seed=0, device='auto',
|
||||||
action_noise_std=0.0, start_timesteps=100, update_style='original',
|
action_noise_std=0.0, learning_starts=100, update_style='original',
|
||||||
create_eval_env=False,
|
create_eval_env=False,
|
||||||
_init_setup_model=True):
|
_init_setup_model=True):
|
||||||
|
|
||||||
super(CEMRL, self).__init__(policy, env, policy_kwargs, verbose,
|
super(CEMRL, self).__init__(policy, env,
|
||||||
buffer_size, learning_rate, seed, device,
|
buffer_size=buffer_size, learning_rate=learning_rate, seed=seed, device=device,
|
||||||
action_noise_std, start_timesteps,
|
action_noise_std=action_noise_std, learning_starts=learning_starts,
|
||||||
policy_freq=policy_freq, batch_size=batch_size,
|
policy_kwargs=policy_kwargs, verbose=verbose,
|
||||||
|
policy_delay=policy_delay, batch_size=batch_size,
|
||||||
create_eval_env=create_eval_env,
|
create_eval_env=create_eval_env,
|
||||||
_init_setup_model=False)
|
_init_setup_model=False)
|
||||||
|
|
||||||
|
|
@ -106,7 +107,7 @@ class CEMRL(TD3):
|
||||||
self.train_critic(replay_data=replay_data)
|
self.train_critic(replay_data=replay_data)
|
||||||
|
|
||||||
# Delayed policy updates
|
# Delayed policy updates
|
||||||
if it % self.policy_freq == 0:
|
if it % self.policy_delay == 0:
|
||||||
self.train_actor(replay_data=replay_data)
|
self.train_actor(replay_data=replay_data)
|
||||||
|
|
||||||
# Get the params back in the population
|
# Get the params back in the population
|
||||||
|
|
@ -134,7 +135,7 @@ class CEMRL(TD3):
|
||||||
episode_reward, episode_timesteps = self.collect_rollouts(self.env, n_episodes=1,
|
episode_reward, episode_timesteps = self.collect_rollouts(self.env, n_episodes=1,
|
||||||
action_noise_std=self.action_noise_std,
|
action_noise_std=self.action_noise_std,
|
||||||
deterministic=False, callback=None,
|
deterministic=False, callback=None,
|
||||||
start_timesteps=self.start_timesteps,
|
learning_starts=self.learning_starts,
|
||||||
num_timesteps=self.num_timesteps,
|
num_timesteps=self.num_timesteps,
|
||||||
replay_buffer=self.replay_buffer)
|
replay_buffer=self.replay_buffer)
|
||||||
episode_num += 1
|
episode_num += 1
|
||||||
|
|
|
||||||
|
|
@ -216,8 +216,8 @@ class BaseRLModel(object):
|
||||||
self.eval_env.seed(seed)
|
self.eval_env.seed(seed)
|
||||||
|
|
||||||
def collect_rollouts(self, env, n_episodes=1, action_noise_std=0.0,
|
def collect_rollouts(self, env, n_episodes=1, action_noise_std=0.0,
|
||||||
deterministic=False, callback=None, remove_timelimits=True,
|
deterministic=False, callback=None,
|
||||||
start_timesteps=0, num_timesteps=0, replay_buffer=None):
|
learning_starts=0, num_timesteps=0, replay_buffer=None):
|
||||||
|
|
||||||
episode_rewards = []
|
episode_rewards = []
|
||||||
total_timesteps = []
|
total_timesteps = []
|
||||||
|
|
@ -231,24 +231,19 @@ class BaseRLModel(object):
|
||||||
episode_reward, episode_timesteps = 0.0, 0
|
episode_reward, episode_timesteps = 0.0, 0
|
||||||
while not done:
|
while not done:
|
||||||
# Select action randomly or according to policy
|
# Select action randomly or according to policy
|
||||||
if num_timesteps < start_timesteps:
|
if num_timesteps < learning_starts:
|
||||||
action = [self.action_space.sample()]
|
action = [self.action_space.sample()]
|
||||||
else:
|
else:
|
||||||
action = self.predict(obs, deterministic=deterministic) / self.max_action
|
action = self.predict(obs, deterministic=deterministic) / self.max_action
|
||||||
|
|
||||||
if action_noise_std > 0:
|
if action_noise_std > 0:
|
||||||
# NOTE: in the original implementation, the noise is applied to the unscaled action
|
# NOTE: in the original implementation of TD3, the noise was applied to the unscaled action
|
||||||
action_noise = np.random.normal(0, action_noise_std, size=self.action_space.shape[0])
|
action_noise = np.random.normal(0, action_noise_std, size=self.action_space.shape[0])
|
||||||
action = (action + action_noise).clip(-1, 1)
|
action = (action + action_noise).clip(-1, 1)
|
||||||
|
|
||||||
# Rescale and perform action
|
# Rescale and perform action
|
||||||
new_obs, reward, done, _ = env.step(self.max_action * action)
|
new_obs, reward, done, _ = env.step(self.max_action * action)
|
||||||
|
|
||||||
# TODO: fix for VecEnv
|
|
||||||
# if hasattr(self.env, '_max_episode_steps') and remove_timelimits:
|
|
||||||
# done_bool = 0 if episode_timesteps + 1 == env._max_episode_steps else float(done)
|
|
||||||
# else:
|
|
||||||
# done_bool = float(done)
|
|
||||||
done_bool = [float(done[0])]
|
done_bool = [float(done[0])]
|
||||||
episode_reward += reward
|
episode_reward += reward
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -27,7 +27,6 @@ class PPO(BaseRLModel):
|
||||||
and https://github.com/ikostrikov/pytorch-a2c-ppo-acktr-gail
|
and https://github.com/ikostrikov/pytorch-a2c-ppo-acktr-gail
|
||||||
and stable_baselines
|
and stable_baselines
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, policy, env, policy_kwargs=None, verbose=0,
|
def __init__(self, policy, env, policy_kwargs=None, verbose=0,
|
||||||
learning_rate=3e-4, seed=0, device='auto',
|
learning_rate=3e-4, seed=0, device='auto',
|
||||||
n_optim=5, batch_size=64, n_steps=256,
|
n_optim=5, batch_size=64, n_steps=256,
|
||||||
|
|
|
||||||
|
|
@ -12,22 +12,50 @@ from torchy_baselines.sac.policies import SACPolicy
|
||||||
|
|
||||||
class SAC(BaseRLModel):
|
class SAC(BaseRLModel):
|
||||||
"""
|
"""
|
||||||
Implementation of Soft Actor-Critic (SAC)
|
Soft Actor-Critic (SAC)
|
||||||
Off-Policy Maximum Entropy Deep Reinforcement Learning with a Stochastic Actor,
|
Off-Policy Maximum Entropy Deep Reinforcement Learning with a Stochastic Actor,
|
||||||
Paper: https://arxiv.org/abs/1801.01290
|
This implementation borrows code from original implementation (https://github.com/haarnoja/sac)
|
||||||
Code: 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
|
||||||
from OpenAI Spinning Up (https://github.com/openai/spinningup) and from the Softlearning repo
|
|
||||||
(https://github.com/rail-berkeley/softlearning/)
|
(https://github.com/rail-berkeley/softlearning/)
|
||||||
|
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
|
||||||
|
|
||||||
Note: we use double q target and not value target as discussed
|
Note: we use double q target and not value target as discussed
|
||||||
in https://github.com/hill-a/stable-baselines/issues/270
|
in https://github.com/hill-a/stable-baselines/issues/270
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, policy, env, policy_kwargs=None, verbose=0,
|
:param policy: (SACPolicy or str) The policy model to use (MlpPolicy, CnnPolicy, ...)
|
||||||
buffer_size=int(1e6), learning_rate=3e-4, seed=0, device='auto',
|
:param env: (Gym environment or str) The environment to learn from (if registered in Gym, can be str)
|
||||||
ent_coef='auto', target_entropy='auto', gamma=0.99,
|
:param learning_rate: (float or callable) learning rate for adam optimizer,
|
||||||
action_noise_std=0.0, start_timesteps=100,
|
the same learning rate will be used for all networks (Q-Values, Actor and Value function)
|
||||||
batch_size=64, create_eval_env=False,
|
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 train_freq: (int) Update the model every `train_freq` steps.
|
||||||
|
: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.
|
||||||
|
:param gradient_steps: (int) How many gradient update after each step
|
||||||
|
:param target_entropy: (str or float) target entropy when learning ent_coef (ent_coef = 'auto')
|
||||||
|
:param action_noise: (ActionNoise) the action noise type (None by default), this can help
|
||||||
|
for hard exploration problem. Cf DDPG for the different action noise type.
|
||||||
|
:param gamma: (float) the discount factor
|
||||||
|
: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 _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, buffer_size=int(1e6),
|
||||||
|
learning_starts=100, train_freq=1, batch_size=64,
|
||||||
|
tau=0.005, ent_coef='auto', target_update_interval=1,
|
||||||
|
gradient_steps=1, target_entropy='auto', action_noise=None,
|
||||||
|
gamma=0.99, action_noise_std=0.0, create_eval_env=False,
|
||||||
|
policy_kwargs=None, verbose=0, seed=0,
|
||||||
_init_setup_model=True):
|
_init_setup_model=True):
|
||||||
|
|
||||||
super(SAC, self).__init__(policy, env, SACPolicy, policy_kwargs, verbose, device,
|
super(SAC, self).__init__(policy, env, SACPolicy, policy_kwargs, verbose, device,
|
||||||
|
|
@ -36,16 +64,24 @@ class SAC(BaseRLModel):
|
||||||
self.max_action = np.abs(self.action_space.high)
|
self.max_action = np.abs(self.action_space.high)
|
||||||
self.action_noise_std = action_noise_std
|
self.action_noise_std = action_noise_std
|
||||||
self.learning_rate = learning_rate
|
self.learning_rate = learning_rate
|
||||||
self.buffer_size = buffer_size
|
|
||||||
self.start_timesteps = start_timesteps
|
|
||||||
self._seed = seed
|
self._seed = seed
|
||||||
self.batch_size = batch_size
|
|
||||||
|
|
||||||
self.ent_coef = ent_coef
|
|
||||||
self.target_entropy = target_entropy
|
self.target_entropy = target_entropy
|
||||||
self.log_ent_coef = None
|
self.log_ent_coef = None
|
||||||
# self.target_update_interval = target_update_interval
|
# self.target_update_interval = target_update_interval
|
||||||
# self.gradient_steps = gradient_steps
|
# self.gradient_steps = gradient_steps
|
||||||
|
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
|
||||||
|
# self.train_freq = train_freq
|
||||||
|
# self.gradient_steps = gradient_steps
|
||||||
|
# self.action_noise = action_noise
|
||||||
self.gamma = gamma
|
self.gamma = gamma
|
||||||
|
|
||||||
if _init_setup_model:
|
if _init_setup_model:
|
||||||
|
|
@ -118,7 +154,7 @@ class SAC(BaseRLModel):
|
||||||
"""
|
"""
|
||||||
return self.max_action * self.select_action(observation)
|
return self.max_action * self.select_action(observation)
|
||||||
|
|
||||||
def train(self, n_iterations, batch_size=64, tau=0.005):
|
def train(self, n_iterations, batch_size=64):
|
||||||
|
|
||||||
for it in range(n_iterations):
|
for it in range(n_iterations):
|
||||||
|
|
||||||
|
|
@ -176,7 +212,7 @@ class SAC(BaseRLModel):
|
||||||
|
|
||||||
# Update target networks
|
# Update target networks
|
||||||
for param, target_param in zip(self.critic.parameters(), self.critic_target.parameters()):
|
for param, target_param in zip(self.critic.parameters(), self.critic_target.parameters()):
|
||||||
target_param.data.copy_(tau * param.data + (1 - tau) * target_param.data)
|
target_param.data.copy_(self.tau * param.data + (1 - self.tau) * target_param.data)
|
||||||
|
|
||||||
def learn(self, total_timesteps, callback=None, log_interval=100,
|
def learn(self, total_timesteps, callback=None, log_interval=100,
|
||||||
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="TD3", reset_num_timesteps=True):
|
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="TD3", reset_num_timesteps=True):
|
||||||
|
|
@ -197,7 +233,7 @@ class SAC(BaseRLModel):
|
||||||
episode_reward, episode_timesteps = self.collect_rollouts(self.env, n_episodes=1,
|
episode_reward, episode_timesteps = self.collect_rollouts(self.env, n_episodes=1,
|
||||||
action_noise_std=self.action_noise_std,
|
action_noise_std=self.action_noise_std,
|
||||||
deterministic=False, callback=None,
|
deterministic=False, callback=None,
|
||||||
start_timesteps=self.start_timesteps,
|
learning_starts=self.learning_starts,
|
||||||
num_timesteps=self.num_timesteps,
|
num_timesteps=self.num_timesteps,
|
||||||
replay_buffer=self.replay_buffer)
|
replay_buffer=self.replay_buffer)
|
||||||
episode_num += 1
|
episode_num += 1
|
||||||
|
|
|
||||||
|
|
@ -12,27 +12,66 @@ from torchy_baselines.td3.policies import TD3Policy
|
||||||
|
|
||||||
class TD3(BaseRLModel):
|
class TD3(BaseRLModel):
|
||||||
"""
|
"""
|
||||||
Implementation of Twin Delayed Deep Deterministic Policy Gradients (TD3)
|
Twin Delayed DDPG (TD3)
|
||||||
|
Addressing Function Approximation Error in Actor-Critic Methods.
|
||||||
|
|
||||||
|
Original implementation: https://github.com/sfujim/TD3
|
||||||
Paper: https://arxiv.org/abs/1802.09477
|
Paper: https://arxiv.org/abs/1802.09477
|
||||||
Code: https://github.com/sfujim/TD3
|
Introduction to TD3: https://spinningup.openai.com/en/latest/algorithms/td3.html
|
||||||
|
|
||||||
|
:param policy: (TD3Policy 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 buffer_size: (int) size of the replay buffer
|
||||||
|
:param learning_rate: (float or callable) learning rate for adam optimizer,
|
||||||
|
the same learning rate will be used for all networks (Q-Values and Actor networks)
|
||||||
|
it can be a function of the current progress (from 1 to 0)
|
||||||
|
:param policy_delay: (int) Policy and target networks will only be updated once every policy_delay steps
|
||||||
|
per training steps. The Q values will be updated policy_delay more often (update every training step).
|
||||||
|
:param learning_starts: (int) how many steps of the model to collect transitions for before learning starts
|
||||||
|
:param gamma: (float) the discount factor
|
||||||
|
:param batch_size: (int) Minibatch size for each gradient update
|
||||||
|
:param train_freq: (int) Update the model every `train_freq` steps.
|
||||||
|
:param gradient_steps: (int) How many gradient update after each step
|
||||||
|
:param tau: (float) the soft update coefficient ("polyak update" of the target networks, between 0 and 1)
|
||||||
|
:param action_noise: (ActionNoise) the action noise type. Cf DDPG for the different action noise type.
|
||||||
|
:param target_policy_noise: (float) Standard deviation of gaussian noise added to target policy
|
||||||
|
(smoothing noise)
|
||||||
|
:param target_noise_clip: (float) Limit for absolute value of target policy smoothing noise.
|
||||||
|
: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 _init_setup_model: (bool) Whether or not to build the network at the creation of the instance
|
||||||
"""
|
"""
|
||||||
def __init__(self, policy, env, policy_kwargs=None, verbose=0,
|
def __init__(self, policy, env, buffer_size=int(1e6), learning_rate=1e-3,
|
||||||
buffer_size=int(1e6), learning_rate=1e-3, seed=0, device='auto',
|
action_noise_std=0.1, policy_delay=2, learning_starts=100,
|
||||||
action_noise_std=0.1, start_timesteps=100, policy_freq=2,
|
gamma=0.99, batch_size=100, train_freq=1000, gradient_steps=1000,
|
||||||
batch_size=100, create_eval_env=False,
|
tau=0.005, action_noise=None, target_policy_noise=0.2, target_noise_clip=0.5,
|
||||||
_init_setup_model=True):
|
create_eval_env=False, policy_kwargs=None, verbose=0,
|
||||||
|
seed=0, device='auto', _init_setup_model=True):
|
||||||
|
|
||||||
super(TD3, self).__init__(policy, env, TD3Policy, policy_kwargs, verbose, device,
|
super(TD3, self).__init__(policy, env, TD3Policy, policy_kwargs, verbose, device,
|
||||||
create_eval_env=create_eval_env)
|
create_eval_env=create_eval_env)
|
||||||
|
|
||||||
self.max_action = np.abs(self.action_space.high)
|
self.max_action = np.abs(self.action_space.high)
|
||||||
self.action_noise_std = action_noise_std
|
self.action_noise_std = action_noise_std
|
||||||
self.learning_rate = learning_rate
|
|
||||||
self.buffer_size = buffer_size
|
self.buffer_size = buffer_size
|
||||||
self.start_timesteps = start_timesteps
|
|
||||||
self._seed = seed
|
self._seed = seed
|
||||||
self.policy_freq = policy_freq
|
|
||||||
|
self.buffer_size = buffer_size
|
||||||
|
# TODO: accept callables
|
||||||
|
self.learning_rate = learning_rate
|
||||||
|
self.learning_starts = learning_starts
|
||||||
|
# self.train_freq = train_freq
|
||||||
|
# self.gradient_steps = gradient_steps
|
||||||
self.batch_size = batch_size
|
self.batch_size = batch_size
|
||||||
|
# self.tau = tau
|
||||||
|
self.gamma = gamma
|
||||||
|
# self.action_noise = action_noise
|
||||||
|
self.policy_delay = policy_delay
|
||||||
|
self.target_noise_clip = target_noise_clip
|
||||||
|
self.target_policy_noise = target_policy_noise
|
||||||
|
|
||||||
if _init_setup_model:
|
if _init_setup_model:
|
||||||
self._setup_model()
|
self._setup_model()
|
||||||
|
|
@ -69,10 +108,11 @@ class TD3(BaseRLModel):
|
||||||
:param deterministic: (bool) Whether or not to return deterministic actions.
|
: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)
|
:return: (np.ndarray, np.ndarray) the model's action and the next state (used in recurrent policies)
|
||||||
"""
|
"""
|
||||||
return self.max_action * self.select_action(observation)
|
# Rescale the action (no need for symmetric action space)
|
||||||
|
return self.action_space.low +\
|
||||||
|
(0.5 * (self.select_action(observation) + 1.0) * (self.action_space.high - self.action_space.low))
|
||||||
|
|
||||||
def train_critic(self, n_iterations=1, batch_size=100, discount=0.99,
|
def train_critic(self, n_iterations=1, batch_size=100, replay_data=None, tau=0.0):
|
||||||
policy_noise=0.2, noise_clip=0.5, replay_data=None, tau=0.0):
|
|
||||||
|
|
||||||
for it in range(n_iterations):
|
for it in range(n_iterations):
|
||||||
# Sample replay buffer
|
# Sample replay buffer
|
||||||
|
|
@ -82,14 +122,14 @@ class TD3(BaseRLModel):
|
||||||
obs, action, next_obs, done, reward = replay_data
|
obs, action, next_obs, done, reward = replay_data
|
||||||
|
|
||||||
# Select action according to policy and add clipped noise
|
# Select action according to policy and add clipped noise
|
||||||
noise = action.clone().data.normal_(0, policy_noise)
|
noise = action.clone().data.normal_(0, self.target_policy_noise)
|
||||||
noise = noise.clamp(-noise_clip, noise_clip)
|
noise = noise.clamp(-self.target_noise_clip, self.target_noise_clip)
|
||||||
next_action = (self.actor_target(next_obs) + noise).clamp(-1, 1)
|
next_action = (self.actor_target(next_obs) + noise).clamp(-1, 1)
|
||||||
|
|
||||||
# Compute the target Q value
|
# Compute the target Q value
|
||||||
target_q1, target_q2 = self.critic_target(next_obs, next_action)
|
target_q1, target_q2 = self.critic_target(next_obs, next_action)
|
||||||
target_q = th.min(target_q1, target_q2)
|
target_q = th.min(target_q1, target_q2)
|
||||||
target_q = reward + ((1 - done) * discount * target_q).detach()
|
target_q = reward + ((1 - done) * self.gamma * target_q).detach()
|
||||||
|
|
||||||
# Get current Q estimates
|
# Get current Q estimates
|
||||||
current_q1, current_q2 = self.critic(obs, action)
|
current_q1, current_q2 = self.critic(obs, action)
|
||||||
|
|
@ -134,7 +174,7 @@ class TD3(BaseRLModel):
|
||||||
for param, target_param in zip(self.actor.parameters(), self.actor_target.parameters()):
|
for param, target_param in zip(self.actor.parameters(), self.actor_target.parameters()):
|
||||||
target_param.data.copy_(tau_actor * param.data + (1 - tau_actor) * target_param.data)
|
target_param.data.copy_(tau_actor * param.data + (1 - tau_actor) * target_param.data)
|
||||||
|
|
||||||
def train(self, n_iterations, batch_size=100, policy_freq=2):
|
def train(self, n_iterations, batch_size=100, policy_delay=2):
|
||||||
|
|
||||||
for it in range(n_iterations):
|
for it in range(n_iterations):
|
||||||
|
|
||||||
|
|
@ -143,7 +183,7 @@ class TD3(BaseRLModel):
|
||||||
self.train_critic(replay_data=replay_data)
|
self.train_critic(replay_data=replay_data)
|
||||||
|
|
||||||
# Delayed policy updates
|
# Delayed policy updates
|
||||||
if it % policy_freq == 0:
|
if it % policy_delay == 0:
|
||||||
self.train_actor(replay_data=replay_data)
|
self.train_actor(replay_data=replay_data)
|
||||||
|
|
||||||
def learn(self, total_timesteps, callback=None, log_interval=100,
|
def learn(self, total_timesteps, callback=None, log_interval=100,
|
||||||
|
|
@ -165,7 +205,7 @@ class TD3(BaseRLModel):
|
||||||
episode_reward, episode_timesteps = self.collect_rollouts(self.env, n_episodes=1,
|
episode_reward, episode_timesteps = self.collect_rollouts(self.env, n_episodes=1,
|
||||||
action_noise_std=self.action_noise_std,
|
action_noise_std=self.action_noise_std,
|
||||||
deterministic=False, callback=None,
|
deterministic=False, callback=None,
|
||||||
start_timesteps=self.start_timesteps,
|
learning_starts=self.learning_starts,
|
||||||
num_timesteps=self.num_timesteps,
|
num_timesteps=self.num_timesteps,
|
||||||
replay_buffer=self.replay_buffer)
|
replay_buffer=self.replay_buffer)
|
||||||
episode_num += 1
|
episode_num += 1
|
||||||
|
|
@ -176,7 +216,7 @@ class TD3(BaseRLModel):
|
||||||
if self.verbose > 1:
|
if self.verbose > 1:
|
||||||
print("Total T: {} Episode Num: {} Episode T: {} Reward: {}".format(
|
print("Total T: {} Episode Num: {} Episode T: {} Reward: {}".format(
|
||||||
self.num_timesteps, episode_num, episode_timesteps, episode_reward))
|
self.num_timesteps, episode_num, episode_timesteps, episode_reward))
|
||||||
self.train(episode_timesteps, batch_size=self.batch_size, policy_freq=self.policy_freq)
|
self.train(episode_timesteps, batch_size=self.batch_size, policy_delay=self.policy_delay)
|
||||||
|
|
||||||
# Evaluate episode
|
# Evaluate episode
|
||||||
if 0 < eval_freq <= timesteps_since_eval and eval_env is not None:
|
if 0 < eval_freq <= timesteps_since_eval and eval_env is not None:
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue