Add tuning script for GAIL (PPO)
This commit is contained in:
@@ -30,7 +30,7 @@ def roll_buffer(buffer, *args, **kwargs):
|
||||
|
||||
def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value,
|
||||
v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma,
|
||||
gae_lambda, delta, backtrack_coeff, backtrack_iters, cg_iters=10, cg_damping=0.1, wasserstein=False, wasserstein_c=None, logger=TerminalLogger()):
|
||||
gae_lambda, delta, backtrack_coeff, backtrack_iters, cg_iters=10, cg_damping=0.1, wasserstein=False, wasserstein_c=None, logger=TerminalLogger(), callback=None, lr_schedulers=[]):
|
||||
|
||||
policy(torch.zeros(env_fn(0).observation_space.shape))
|
||||
policy = ReparamPolicy(policy)
|
||||
@@ -39,10 +39,15 @@ def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value
|
||||
logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0])
|
||||
|
||||
for epoch in tqdm(range(epochs)):
|
||||
generator_data = Buffer(*rollout(env_fn, policy, rollout_episodes, rollout_steps))
|
||||
states, actions, rewards, dones, collisions = rollout(env_fn, policy, rollout_episodes, rollout_steps)
|
||||
generator_data = Buffer(states, actions, rewards, dones)
|
||||
|
||||
logger.add_scalar('gen/mean_episode_length', (~generator_data.dones).sum() / generator_data.states.shape[0], epoch)
|
||||
logger.add_scalar('gen/mean_reward_per_episode', generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0], epoch)
|
||||
gen_mean_episode_length = (~generator_data.dones).sum() / generator_data.states.shape[0]
|
||||
logger.add_scalar('gen/mean_episode_length', gen_mean_episode_length, epoch)
|
||||
gen_mean_reward_per_episode = generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0]
|
||||
logger.add_scalar('gen/mean_reward_per_episode', gen_mean_reward_per_episode, epoch)
|
||||
gen_collision_rate = (1. * collisions.any(-1)).mean()
|
||||
logger.add_scalar('gen/collision_rate', gen_collision_rate, epoch)
|
||||
|
||||
discriminator, loss = train_discriminator(expert_data, generator_data, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c)
|
||||
if wasserstein:
|
||||
@@ -50,25 +55,45 @@ def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value
|
||||
else:
|
||||
generator_data.rewards = -F.logsigmoid(discriminator(generator_data.states, generator_data.actions))
|
||||
logger.add_scalar('disc/final_loss', loss, epoch)
|
||||
logger.add_scalar('disc/mean_reward_per_episode', generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0], epoch)
|
||||
disc_mean_reward_per_episode = generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0]
|
||||
logger.add_scalar('disc/mean_reward_per_episode', disc_mean_reward_per_episode, epoch)
|
||||
|
||||
value, policy = trpo_step(value, policy, generator_data.states, generator_data.actions, generator_data.rewards, generator_data.dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters, cg_damping)
|
||||
expert_data = roll_buffer(expert_data, shifts=-3, dims=0)
|
||||
|
||||
if callback is not None:
|
||||
callback({
|
||||
'epoch': epoch,
|
||||
'value': value,
|
||||
'policy': policy,
|
||||
'gen/mean_episode_length': gen_mean_episode_length.item(),
|
||||
'gen/mean_reward_per_episode': gen_mean_reward_per_episode.item(),
|
||||
'gen/collision_rate': gen_collision_rate.item(),
|
||||
'disc/mean_reward_per_episode': disc_mean_reward_per_episode.item(),
|
||||
})
|
||||
|
||||
for lr_scheduler in lr_schedulers:
|
||||
lr_scheduler.step()
|
||||
|
||||
return value, policy
|
||||
|
||||
def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value,
|
||||
v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma,
|
||||
gae_lambda, clip_ratio, pi_opt, pi_iters, target_kl=None, max_grad_norm=None, wasserstein=False, wasserstein_c=None, logger=TerminalLogger()):
|
||||
gae_lambda, clip_ratio, pi_opt, pi_iters, target_kl=None, max_grad_norm=None, wasserstein=False, wasserstein_c=None, logger=TerminalLogger(), callback=None, lr_schedulers=[]):
|
||||
|
||||
logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0])
|
||||
logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0])
|
||||
|
||||
for epoch in range(epochs):
|
||||
generator_data = Buffer(*rollout(env_fn, policy, rollout_episodes, rollout_steps))
|
||||
states, actions, rewards, dones, collisions = rollout(env_fn, policy, rollout_episodes, rollout_steps)
|
||||
generator_data = Buffer(states, actions, rewards, dones)
|
||||
|
||||
logger.add_scalar('gen/mean_episode_length', (~generator_data.dones).sum() / generator_data.states.shape[0], epoch)
|
||||
logger.add_scalar('gen/mean_reward_per_episode', generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0], epoch)
|
||||
gen_mean_episode_length = (~generator_data.dones).sum() / generator_data.states.shape[0]
|
||||
logger.add_scalar('gen/mean_episode_length', gen_mean_episode_length, epoch)
|
||||
gen_mean_reward_per_episode = generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0]
|
||||
logger.add_scalar('gen/mean_reward_per_episode', gen_mean_reward_per_episode, epoch)
|
||||
gen_collision_rate = (1. * collisions.any(-1)).mean()
|
||||
logger.add_scalar('gen/collision_rate', gen_collision_rate, epoch)
|
||||
|
||||
discriminator, loss = train_discriminator(expert_data, generator_data, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c)
|
||||
if wasserstein:
|
||||
@@ -76,10 +101,25 @@ def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, v
|
||||
else:
|
||||
generator_data.rewards = -F.logsigmoid(discriminator(generator_data.states, generator_data.actions))
|
||||
logger.add_scalar('disc/final_loss', loss, epoch)
|
||||
logger.add_scalar('disc/mean_reward_per_episode', generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0], epoch)
|
||||
disc_mean_reward_per_episode = generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0]
|
||||
logger.add_scalar('disc/mean_reward_per_episode', disc_mean_reward_per_episode, epoch)
|
||||
|
||||
value, policy = ppo_step(value, policy, generator_data.states, generator_data.actions, generator_data.rewards, generator_data.dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm)
|
||||
expert_data = roll_buffer(expert_data, shifts=-3, dims=0)
|
||||
|
||||
if callback is not None:
|
||||
callback({
|
||||
'epoch': epoch,
|
||||
'value': value,
|
||||
'policy': policy,
|
||||
'gen/mean_episode_length': gen_mean_episode_length.item(),
|
||||
'gen/mean_reward_per_episode': gen_mean_reward_per_episode.item(),
|
||||
'gen/collision_rate': gen_collision_rate.item(),
|
||||
'disc/mean_reward_per_episode': disc_mean_reward_per_episode.item(),
|
||||
})
|
||||
|
||||
for lr_scheduler in lr_schedulers:
|
||||
lr_scheduler.step()
|
||||
|
||||
return value, policy
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@ import torch
|
||||
import gym
|
||||
from stable_baselines3.common.vec_env import DummyVecEnv as VecEnv
|
||||
from tqdm import tqdm
|
||||
import numpy as np
|
||||
|
||||
def rollout(env_fn, policy, n_episodes, max_steps_per_episode):
|
||||
env = env_fn(0)
|
||||
@@ -9,23 +10,27 @@ def rollout(env_fn, policy, n_episodes, max_steps_per_episode):
|
||||
actions = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.action_space.shape)
|
||||
rewards = torch.zeros(n_episodes, max_steps_per_episode + 1)
|
||||
dones = torch.ones(n_episodes, max_steps_per_episode + 1, dtype=bool)
|
||||
collisions = torch.zeros(n_episodes, max_steps_per_episode, dtype=bool)
|
||||
|
||||
env = VecEnv(list(map(lambda i: (lambda: env_fn(i)), range(n_episodes))))
|
||||
|
||||
states[:, 0] = torch.tensor(env.reset()).clone().detach()
|
||||
dones[:, 0] = False
|
||||
|
||||
for s in range(max_steps_per_episode):
|
||||
for s in tqdm(range(max_steps_per_episode), 'Rollout'):
|
||||
actions[:, s] = policy.sample(policy(states[:, s])).clone().detach()
|
||||
|
||||
clipped_actions = actions[:, s]
|
||||
if isinstance(env.action_space, gym.spaces.Box):
|
||||
clipped_actions = torch.clamp(clipped_actions, torch.from_numpy(env.action_space.low), torch.from_numpy(env.action_space.high))
|
||||
|
||||
o, r, d, _ = env.step(clipped_actions)
|
||||
o, r, d, info = env.step(clipped_actions)
|
||||
states[:, s + 1] = torch.tensor(o).clone().detach()
|
||||
rewards[:, s] = torch.tensor(r).clone().detach()
|
||||
dones[:, s + 1] = torch.tensor(d).clone().detach()
|
||||
collisions[:, s] = torch.from_numpy(np.stack([
|
||||
i['collision'] for i in info
|
||||
])).detach().clone()
|
||||
|
||||
dones = dones.cumsum(1) > 0
|
||||
|
||||
@@ -34,7 +39,7 @@ def rollout(env_fn, policy, n_episodes, max_steps_per_episode):
|
||||
rewards = rewards[:, :max_steps_per_episode]
|
||||
dones = dones[:, :max_steps_per_episode]
|
||||
|
||||
return states, actions, rewards, dones
|
||||
return states, actions, rewards, dones, collisions
|
||||
|
||||
|
||||
def rollout_sb3(env, policy, n_episodes, max_steps_per_episode):
|
||||
|
||||
Reference in New Issue
Block a user