Merge branch 'main' into idm_upgrade
This commit is contained in:
@@ -18,23 +18,17 @@ class Discriminator(nn.Module):
|
||||
|
||||
class DeepsetDiscriminator(nn.Module):
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self, n_hidden_layers_element=3, n_hidden_layers_global=2, hidden_layer_size=10, activation=nn.Tanh):
|
||||
super().__init__()
|
||||
self.elem = nn.Sequential(
|
||||
nn.LazyLinear(10),
|
||||
nn.Tanh(),
|
||||
nn.LazyLinear(10),
|
||||
nn.Tanh(),
|
||||
nn.LazyLinear(10),
|
||||
)
|
||||
self.glob = nn.Sequential(
|
||||
nn.LazyLinear(10),
|
||||
nn.Tanh(),
|
||||
nn.LazyLinear(10),
|
||||
nn.Tanh(),
|
||||
nn.LazyLinear(1),
|
||||
)
|
||||
|
||||
|
||||
layers_elem = sum([[nn.LazyLinear(hidden_layer_size),
|
||||
activation()] for _ in range(n_hidden_layers_element)], [])
|
||||
self.elem = nn.Sequential(*layers_elem)
|
||||
|
||||
layers_glob = sum([[nn.LazyLinear(hidden_layer_size),
|
||||
activation()] for _ in range(n_hidden_layers_global)], [])
|
||||
self.glob = nn.Sequential(*layers_glob, nn.LazyLinear(1))
|
||||
|
||||
def forward(self, states, actions):
|
||||
actions = actions.unsqueeze(-2)
|
||||
actions = actions.expand(*actions.shape[:-2], states.shape[-2], actions.shape[-1])
|
||||
|
||||
@@ -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)
|
||||
@@ -38,11 +38,16 @@ def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value
|
||||
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 tqdm(range(epochs)):
|
||||
generator_data = Buffer(*rollout(env_fn, policy, rollout_episodes, rollout_steps))
|
||||
for epoch in range(epochs):
|
||||
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
|
||||
|
||||
|
||||
@@ -36,30 +36,31 @@ class BasePolicy(nn.Module):
|
||||
|
||||
class Policy(BasePolicy):
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
def __init__(self, *args, hidden_layer_size=50, n_hidden_layers=2, activation=nn.Tanh, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.nn = nn.Sequential(
|
||||
nn.LazyLinear(50),
|
||||
nn.Tanh(),
|
||||
nn.LazyLinear(50),
|
||||
nn.Tanh(),
|
||||
nn.LazyLinear(2 * self.action_dim),
|
||||
)
|
||||
layers = sum([[nn.LazyLinear(hidden_layer_size),
|
||||
activation()] for _ in range(n_hidden_layers)],[])
|
||||
self.nn = nn.Sequential(*layers, nn.LazyLinear(2 *self.action_dim))
|
||||
|
||||
# old
|
||||
# self.nn = nn.Sequential(
|
||||
# nn.LazyLinear(50),
|
||||
# nn.Tanh(),
|
||||
# nn.LazyLinear(50),
|
||||
# nn.Tanh(),
|
||||
# nn.LazyLinear(2 * self.action_dim),
|
||||
#)
|
||||
|
||||
def forward(self, states):
|
||||
return self.nn(states)
|
||||
|
||||
class DiscretePolicy(BasePolicy):
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
def __init__(self, *args, hidden_layer_size=50, n_hidden_layers=2, activation=nn.Tanh, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.nn = nn.Sequential(
|
||||
nn.LazyLinear(50),
|
||||
nn.Tanh(),
|
||||
nn.LazyLinear(50),
|
||||
nn.Tanh(),
|
||||
nn.LazyLinear(self.action_dim),
|
||||
)
|
||||
layers = sum([[nn.LazyLinear(hidden_layer_size),
|
||||
activation()] for _ in range(n_hidden_layers)],[])
|
||||
self.nn = nn.Sequential(*layers, nn.LazyLinear(self.action_dim))
|
||||
|
||||
def forward(self, states):
|
||||
return self.nn(states)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -14,7 +14,8 @@ from src.core.reparam_module import ReparamPolicy, ReparamSafePolicy
|
||||
from src.options import envs as options_envs2
|
||||
from src.safe_options.policy import SetMaskedDiscretePolicy
|
||||
from src.safe_options import options as options_envs3
|
||||
|
||||
from src.util.wrappers import IntersimpleTimeLimit
|
||||
import os
|
||||
from typing import Optional, List, Dict, Tuple
|
||||
import torch
|
||||
import numpy as np
|
||||
@@ -36,46 +37,47 @@ def load_policy(method:str,
|
||||
Returns:
|
||||
policy (Optional[BaseAlgorithm]): the policy to evaluate
|
||||
"""
|
||||
ml = torch.device('cpu') if not torch.cuda.is_available() else None
|
||||
if method == 'idm':
|
||||
policy = IDMRulePolicy(env, **policy_kwargs)
|
||||
elif method == 'bc':
|
||||
policy = SetPolicy(env.action_space.shape[-1])
|
||||
policy.load_state_dict(torch.load(policy_file))
|
||||
policy.load_state_dict(torch.load(policy_file, map_location=ml))
|
||||
policy.eval()
|
||||
elif method == 'gail-trpo':
|
||||
policy = SetPolicy(env.action_space.shape[-1])
|
||||
policy(torch.zeros(env.observation_space.shape))
|
||||
policy = ReparamPolicy(policy)
|
||||
policy.load_state_dict(torch.load(policy_file, map_location=ml))
|
||||
policy.eval()
|
||||
elif method == 'gail':
|
||||
policy = SetPolicy(env.action_space.shape[-1])
|
||||
policy(torch.zeros(env.observation_space.shape))
|
||||
policy = ReparamPolicy(policy)
|
||||
policy.load_state_dict(torch.load(policy_file))
|
||||
policy.eval()
|
||||
elif method == 'gail-ppo':
|
||||
policy = SetPolicy(env.action_space.shape[-1])
|
||||
policy.load_state_dict(torch.load(policy_file))
|
||||
policy.load_state_dict(torch.load(policy_file, map_location=ml))
|
||||
policy.eval()
|
||||
elif method == 'rail':
|
||||
raise NotImplementedError
|
||||
elif method == 'ogail':
|
||||
elif method == 'hail-trpo':
|
||||
policy = SetDiscretePolicy(env.action_space.n)
|
||||
policy(torch.zeros(env.observation_space.shape))
|
||||
policy = ReparamPolicy(policy)
|
||||
policy.load_state_dict(torch.load(policy_file))
|
||||
policy.load_state_dict(torch.load(policy_file, map_location=ml))
|
||||
policy.eval()
|
||||
elif method == 'ogail-ppo':
|
||||
elif method == 'hail':
|
||||
policy = SetDiscretePolicy(env.action_space.n)
|
||||
policy.load_state_dict(torch.load(policy_file))
|
||||
policy.load_state_dict(torch.load(policy_file, map_location=ml))
|
||||
policy.eval()
|
||||
elif method == 'sgail':
|
||||
elif method == 'shail-trpo':
|
||||
policy = SetMaskedDiscretePolicy(env.action_space.n)
|
||||
policy(
|
||||
torch.zeros(env.observation_space['observation'].shape),
|
||||
torch.zeros(env.observation_space['safe_actions'].shape)
|
||||
)
|
||||
policy = ReparamSafePolicy(policy)
|
||||
policy.load_state_dict(torch.load(policy_file))
|
||||
policy.load_state_dict(torch.load(policy_file, map_location=ml))
|
||||
policy.eval()
|
||||
elif method == 'sgail-ppo':
|
||||
elif method == 'shail':
|
||||
policy = SetMaskedDiscretePolicy(env.action_space.n)
|
||||
policy.load_state_dict(torch.load(policy_file))
|
||||
policy.load_state_dict(torch.load(policy_file, map_location=ml))
|
||||
policy.eval()
|
||||
else:
|
||||
raise NotImplementedError
|
||||
@@ -201,7 +203,7 @@ def evaluate_policy(locations:List[Tuple[int,int]],
|
||||
|
||||
# iterate through vehicles
|
||||
for i, location in tqdm(enumerate(locations)):
|
||||
|
||||
|
||||
# add roundabout and track to environent
|
||||
iround, track = location
|
||||
rname = intersim.LOCATIONS[iround]
|
||||
@@ -214,7 +216,14 @@ def evaluate_policy(locations:List[Tuple[int,int]],
|
||||
|
||||
# initialize environment
|
||||
Env = envs_dict[env_class]
|
||||
eval_env = Env(**env_kwargs)
|
||||
|
||||
# wrap in TimeLimit
|
||||
if 'max_episode_steps' in it_env_kwargs.keys():
|
||||
steps = it_env_kwargs.pop('max_episode_steps')
|
||||
eval_env = IntersimpleTimeLimit(Env(**it_env_kwargs),
|
||||
max_episode_steps=steps)
|
||||
else:
|
||||
eval_env = Env(**it_env_kwargs)
|
||||
evaluator = IntersimpleEvaluation(eval_env)
|
||||
|
||||
# load policy
|
||||
@@ -238,6 +247,14 @@ def summary_metrics(metrics:List[Dict[str,list]]) -> Dict[str,float]:
|
||||
"""
|
||||
# keys = ['col_all','v_all', 'a_all','j_all', 'v_avg', 'a_avg', 'col', 'brake', 't']
|
||||
summary_metrics = {}
|
||||
|
||||
# mean travel distance
|
||||
dt = 0.1
|
||||
travel_ds = []
|
||||
for iRound in range(len(metrics)):
|
||||
for iTraj in range(len(metrics[iRound]['v_all'])):
|
||||
travel_ds.append(dt*sum(metrics[iRound]['v_all'][iTraj]))
|
||||
summary_metrics['mean travel distance'] = sum(travel_ds)/len(travel_ds)
|
||||
|
||||
# average average-velocity
|
||||
all_vavgs = sum([d['v_avg'] for d in metrics],[]) # aggregate to single list
|
||||
@@ -265,6 +282,7 @@ def summary_metrics(metrics:List[Dict[str,list]]) -> Dict[str,float]:
|
||||
# collision rate
|
||||
all_collisions = sum([d['col'] for d in metrics],[]) # aggregate to single list
|
||||
summary_metrics['collision rate'] = sum(all_collisions)/len(all_collisions)
|
||||
summary_metrics['success rate'] = 1 - summary_metrics['collision rate']
|
||||
|
||||
# hard brake rate
|
||||
all_hard_brakes = sum([d['brake'] for d in metrics],[]) # aggregate to single list
|
||||
@@ -273,6 +291,8 @@ def summary_metrics(metrics:List[Dict[str,list]]) -> Dict[str,float]:
|
||||
# average number of timesteps
|
||||
all_ts = sum([d['t'] for d in metrics],[]) # aggregate to single list
|
||||
summary_metrics['mean episode length'] = sum(all_ts)/len(all_ts)
|
||||
summary_metrics['mean episode time'] = summary_metrics['mean episode length'] * dt
|
||||
|
||||
|
||||
for key in summary_metrics.keys():
|
||||
print(f'{key}: {summary_metrics[key]}')
|
||||
@@ -303,7 +323,7 @@ def comparison_metrics(policy_metrics:List[Dict[str,list]],
|
||||
expert_traj.append(np.vstack((expert_metrics[iR]['x_all'][iTraj], expert_metrics[iR]['y_all'][iTraj])))
|
||||
policy_traj.append(np.vstack((policy_metrics[iR]['x_all'][iTraj], policy_metrics[iR]['y_all'][iTraj])))
|
||||
assert len(expert_traj)==len(policy_traj)
|
||||
comparison_metrics['rwse'] = rwse(expert_traj, policy_traj)
|
||||
comparison_metrics.update(rwse(expert_traj, policy_traj))
|
||||
|
||||
# average velocity shortfall
|
||||
expert_vavg = np.array(sum([d['v_avg'] for d in expert_metrics],[]))
|
||||
@@ -355,14 +375,28 @@ def eval_main(
|
||||
policy_file (str): path to saved policy
|
||||
env (str): environment class
|
||||
method (str): method (expert, bc, gail, rail, hgail, hrail)
|
||||
|
||||
Returns:
|
||||
outbase (str): string to outbase
|
||||
"""
|
||||
print(f'Evaluating {method} on {env}')
|
||||
print(f'#############################################################################')
|
||||
print(f'Evaluating {method} from file {policy_file} on {env} at locations {locations}')
|
||||
print(f'#############################################################################')
|
||||
|
||||
# set seed
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
pfilename = policy_file.split('/')[-1].split('.')[0]
|
||||
outbase = f'out/{method}/{pfilename}_seed{seed}'
|
||||
locstr = 'loc_'+'_'.join([f'r{ro}t{tr}' for (ro,tr) in locations])
|
||||
if policy_file == '':
|
||||
method_path = method
|
||||
name_base = method
|
||||
else:
|
||||
path_items = policy_file.split('/')
|
||||
name_base = path_items[-1].split('.')[0]
|
||||
method_path = ('/').join(path_items[1:-1])
|
||||
outfolder = os.path.join('out',method_path,locstr)
|
||||
filebase = name_base + f'_tseed{seed}'
|
||||
outbase = os.path.join(outfolder,filebase)
|
||||
|
||||
# load expert metrics
|
||||
expert_metrics = generate_expert_metrics(locations)
|
||||
@@ -381,7 +415,9 @@ def eval_main(
|
||||
save_metrics(smetrics, outbase+'_summary.pkl')
|
||||
cmetrics = comparison_metrics(policy_metrics, expert_metrics, outbase=outbase)
|
||||
save_metrics(cmetrics, outbase+'_comparison.pkl')
|
||||
|
||||
return outbase
|
||||
|
||||
if __name__=='__main__':
|
||||
import fire
|
||||
fire.Fire(eval_main)
|
||||
fire.Fire(eval_main)
|
||||
|
||||
@@ -6,8 +6,9 @@ from typing import Callable, Dict, Optional
|
||||
import os
|
||||
import pickle
|
||||
from tqdm import tqdm
|
||||
from src.util.wrappers import IntersimpleTimeLimit
|
||||
from src.options.envs import OptionsEnv
|
||||
from src.util.wrappers import OptionsTimeLimit
|
||||
from src.safe_options.options import SafeOptionsEnv
|
||||
|
||||
class IntersimpleEvaluation:
|
||||
"""
|
||||
@@ -36,7 +37,10 @@ class IntersimpleEvaluation:
|
||||
self.env = eval_env
|
||||
self.n_episodes = eval_env.nv
|
||||
self.use_pbar = use_pbar
|
||||
self.is_options_env = isinstance(self.env, (OptionsEnv, OptionsTimeLimit))
|
||||
if isinstance(self.env, IntersimpleTimeLimit):
|
||||
self.is_options_env = isinstance(self.env.env, (OptionsEnv, SafeOptionsEnv))
|
||||
else:
|
||||
self.is_options_env = isinstance(self.env, (OptionsEnv, SafeOptionsEnv))
|
||||
|
||||
# metrics present on every step of every episode
|
||||
self.metric_keys_all = ['x_all', 'y_all', 'v_all', 'a_all', 'col_all']
|
||||
@@ -135,8 +139,8 @@ class IntersimpleEvaluation:
|
||||
self._metrics['a_all'][_agent].append(info['action_taken'][_agent,0].item())
|
||||
col = info['collision']
|
||||
|
||||
if col:
|
||||
assert done
|
||||
# if col: # commenting out if we dont want to end on collision
|
||||
# assert done
|
||||
self._metrics['col_all'][_agent].append(col)
|
||||
|
||||
if done and self.use_pbar:
|
||||
|
||||
@@ -4,10 +4,58 @@ import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
from torch.utils.data import DataLoader
|
||||
from intersim import collisions
|
||||
from typing import List
|
||||
from typing import List, Dict
|
||||
# import tikzplotlib
|
||||
|
||||
def rwse(expert:List[np.ndarray], policy:List[np.ndarray], dt:float=0.1) -> float:
|
||||
def rwse(expert:List[np.ndarray], policy:List[np.ndarray], dt:float=0.1) -> Dict[str,float]:
|
||||
"""
|
||||
Calculate average mean squared displacement error
|
||||
|
||||
Args:
|
||||
expert (List[np.ndarray]): all position trajectories for all expert rollouts
|
||||
policy (List[np.ndarray]): all position trajectories for all policy rollouts
|
||||
|
||||
each trajectory in the list should have shape (2, T). however expert[i] might have a
|
||||
different T than policy[i]
|
||||
|
||||
Returns
|
||||
rwse_dict (Dict[str,float]): dict of different RWSEs
|
||||
"""
|
||||
assert len(expert) == len(policy)
|
||||
|
||||
# calculate rwse
|
||||
times = [1,2,5,10,15,20,25,30]
|
||||
time_indices = [int(t/dt) for t in times]
|
||||
rwse_dict_keys = [f'rwse_{t}s' for t in times]+['rwse_end']
|
||||
se_dict = {key:[] for key in rwse_dict_keys}
|
||||
for expert_trajectory, policy_trajectory in zip(expert, policy):
|
||||
_, T1 = expert_trajectory.shape
|
||||
_, T2 = policy_trajectory.shape
|
||||
minT = min(T1, T2)
|
||||
|
||||
crop_expert_trajectory = expert_trajectory[:, :minT]
|
||||
crop_policy_trajectory = policy_trajectory[:, :minT]
|
||||
|
||||
# square error along every time
|
||||
se = ((crop_policy_trajectory - crop_expert_trajectory)**2).sum(0)
|
||||
|
||||
# add to dict with appropriate indexing
|
||||
for time, idx in zip(times, time_indices):
|
||||
if minT >= idx:
|
||||
se_dict[f'rwse_{time}s'].append(se[idx-1])
|
||||
se_dict['rwse_end'].append(se[-1])
|
||||
|
||||
assert len(se_dict['rwse_end']) == len(expert)
|
||||
|
||||
# print how many trajectories of each time:
|
||||
for key in rwse_dict_keys:
|
||||
print('%s has %i elements'%(key, len(se_dict[key])))
|
||||
|
||||
rwse_dict = {key:np.mean(np.array(se_dict[key]))**0.5 for key in rwse_dict_keys}
|
||||
|
||||
return rwse_dict
|
||||
|
||||
def rwse_basic(expert:List[np.ndarray], policy:List[np.ndarray], dt:float=0.1) -> float:
|
||||
"""
|
||||
Calculate average mean squared displacement error
|
||||
|
||||
@@ -42,8 +90,6 @@ def rwse(expert:List[np.ndarray], policy:List[np.ndarray], dt:float=0.1) -> floa
|
||||
|
||||
return avg_rwse
|
||||
|
||||
|
||||
|
||||
def visualize_distribution(expert, policy, filestr):
|
||||
"""
|
||||
Visualize two distributions
|
||||
|
||||
@@ -2,6 +2,7 @@ import pickle
|
||||
import os
|
||||
import numpy as np
|
||||
from typing import List,Dict
|
||||
|
||||
def save_metrics(metrics:dict, filestr:str):
|
||||
"""
|
||||
Save metric dict to filestr
|
||||
@@ -33,13 +34,20 @@ def load_metrics(filestr:str):
|
||||
metrics = pickle.load(f)
|
||||
return metrics
|
||||
|
||||
def average_metrics(metric_list:List[Dict[str,float]]):
|
||||
def average_metrics(metric_list:List[Dict[str,float]], verbose:bool=True) ->Dict[str, tuple]:
|
||||
"""
|
||||
Average all the metrics in the list
|
||||
|
||||
Args:
|
||||
metric_list (list of dicts): list of metric dicts which each map a string to a float
|
||||
verbose (bool): whether to print avg metrics
|
||||
Returns:
|
||||
average_metrics (Dict[str, tuple])
|
||||
"""
|
||||
average_metrics = {}
|
||||
if len(metric_list) == 0:
|
||||
return average_metrics
|
||||
|
||||
keys = list(metric_list[0].keys())
|
||||
N = len(metric_list)
|
||||
master_dict = {key:[] for key in keys}
|
||||
@@ -47,31 +55,41 @@ def average_metrics(metric_list:List[Dict[str,float]]):
|
||||
for i in range(N):
|
||||
master_dict[key].append(metric_list[i][key])
|
||||
master_dict[key] = np.array(master_dict[key])
|
||||
mu = np.mean(master_dict[key])
|
||||
std2 = np.std(master_dict[key])*2
|
||||
print(f'{key}: {mu} \pm {std2}')
|
||||
mu = np.nanmean(master_dict[key])
|
||||
std2 = np.nanstd(master_dict[key])*2
|
||||
if verbose:
|
||||
print(f'{key}: {mu} \pm {std2}')
|
||||
average_metrics[key] = (mu, std2)
|
||||
return average_metrics
|
||||
|
||||
|
||||
def load_and_average(path:str):
|
||||
def load_and_average(path:str, verbose:bool=True):
|
||||
"""
|
||||
Load and average all metric files in a particular folder
|
||||
|
||||
Args:
|
||||
path (str)
|
||||
verbose (bool): whether to print avg metrics
|
||||
Returns:
|
||||
avg_metrics (Dict[str, tuple])
|
||||
"""
|
||||
assert os.path.isdir(path)
|
||||
|
||||
# summary metrics
|
||||
summary_files = [os.path.join(path,f) for f in os.listdir(path) if f.endswith('summary.pkl')]
|
||||
print(*summary_files, sep='\n')
|
||||
if verbose:
|
||||
print(*summary_files, sep='\n')
|
||||
all_summary_metrics = [load_metrics(f) for f in summary_files]
|
||||
average_metrics(all_summary_metrics)
|
||||
avg_metrics = average_metrics(all_summary_metrics, verbose=verbose)
|
||||
|
||||
# comparison metrics
|
||||
comp_files = [os.path.join(path,f) for f in os.listdir(path) if f.endswith('comparison.pkl')]
|
||||
print(*comp_files, sep='\n')
|
||||
if verbose:
|
||||
print(*comp_files, sep='\n')
|
||||
all_comp_metrics = [load_metrics(f) for f in comp_files]
|
||||
average_metrics(all_comp_metrics)
|
||||
comp_avg = average_metrics(all_comp_metrics, verbose=verbose)
|
||||
avg_metrics.update(comp_avg)
|
||||
return avg_metrics
|
||||
|
||||
if __name__=='__main__':
|
||||
import fire
|
||||
|
||||
@@ -14,7 +14,7 @@ from src.options.envs import OptionsEnv
|
||||
from src.safe_options.collisions import feasible
|
||||
|
||||
from intersim.envs import IntersimpleLidarFlatIncrementingAgent
|
||||
from src.util.wrappers import OptionsTimeLimit, Setobs, TransformObservation
|
||||
from src.util.wrappers import Setobs, TransformObservation
|
||||
|
||||
@dataclass
|
||||
class Buffer:
|
||||
@@ -47,14 +47,18 @@ 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)):
|
||||
hl_data, ll_data = rollout(env_fn, policy, rollout_episodes, rollout_steps)
|
||||
hl_data, ll_data, collisions = rollout(env_fn, policy, rollout_episodes, rollout_steps)
|
||||
generator_data = OptionsRollout(HLBuffer(*hl_data), Buffer(*ll_data))
|
||||
|
||||
generator_data.ll.actions += 0.1 * torch.randn_like(generator_data.ll.actions)
|
||||
|
||||
logger.add_scalar('gen/mean_episode_length', (~generator_data.ll.dones).sum() / generator_data.ll.states.shape[0], epoch)
|
||||
logger.add_scalar('gen/mean_reward_per_episode', generator_data.hl.rewards[~generator_data.hl.dones].sum() / generator_data.hl.states.shape[0], epoch)
|
||||
gen_mean_episode_length = (~generator_data.ll.dones).sum() / generator_data.ll.states.shape[0]
|
||||
logger.add_scalar('gen/mean_episode_length', gen_mean_episode_length , epoch)
|
||||
gen_mean_reward_per_episode = generator_data.hl.rewards[~generator_data.hl.dones].sum() / generator_data.hl.states.shape[0]
|
||||
logger.add_scalar('gen/mean_reward_per_episode', gen_mean_reward_per_episode, epoch)
|
||||
logger.add_scalar('gen/unsafe_probability_mass', policy.unsafe_probability_mass(policy(generator_data.hl.states[~generator_data.hl.dones], generator_data.hl.safe_actions[~generator_data.hl.dones])).mean(), 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.ll, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c)
|
||||
if wasserstein:
|
||||
@@ -62,7 +66,8 @@ def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value
|
||||
else:
|
||||
generator_data.ll.rewards = -F.logsigmoid(discriminator(generator_data.ll.states, generator_data.ll.actions))
|
||||
logger.add_scalar('disc/final_loss', loss, epoch)
|
||||
logger.add_scalar('disc/mean_reward_per_episode', generator_data.ll.rewards[~generator_data.ll.dones].sum() / generator_data.ll.states.shape[0], epoch)
|
||||
disc_mean_reward_per_episode = generator_data.ll.rewards[~generator_data.ll.dones].sum() / generator_data.ll.states.shape[0]
|
||||
logger.add_scalar('disc/mean_reward_per_episode', disc_mean_reward_per_episode , epoch)
|
||||
|
||||
#assert generator_data.ll.rewards.shape == generator_data.ll.dones.shape
|
||||
generator_data.hl.rewards = torch.where(~generator_data.ll.dones, generator_data.ll.rewards, torch.tensor(0.)).sum(-1)
|
||||
@@ -71,26 +76,37 @@ def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value
|
||||
expert_data = roll_buffer(expert_data, shifts=-3, dims=0)
|
||||
|
||||
if callback is not None:
|
||||
callback(epoch, value, policy)
|
||||
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(),
|
||||
})
|
||||
|
||||
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(), callback=None):
|
||||
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):
|
||||
hl_data, ll_data = rollout(env_fn, policy, rollout_episodes, rollout_steps)
|
||||
hl_data, ll_data, collisions = rollout(env_fn, policy, rollout_episodes, rollout_steps)
|
||||
generator_data = OptionsRollout(HLBuffer(*hl_data), Buffer(*ll_data))
|
||||
|
||||
generator_data.ll.actions += 0.1 * torch.randn_like(generator_data.ll.actions)
|
||||
|
||||
logger.add_scalar('gen/mean_episode_length', (~generator_data.ll.dones).sum() / generator_data.ll.states.shape[0], epoch)
|
||||
logger.add_scalar('gen/mean_reward_per_episode', generator_data.hl.rewards[~generator_data.hl.dones].sum() / generator_data.hl.states.shape[0], epoch)
|
||||
gen_mean_episode_length = (~generator_data.ll.dones).sum() / generator_data.ll.states.shape[0]
|
||||
logger.add_scalar('gen/mean_episode_length', gen_mean_episode_length, epoch)
|
||||
gen_mean_reward_per_episode = generator_data.hl.rewards[~generator_data.hl.dones].sum() / generator_data.hl.states.shape[0]
|
||||
logger.add_scalar('gen/mean_reward_per_episode', gen_mean_reward_per_episode, epoch)
|
||||
logger.add_scalar('gen/unsafe_probability_mass', policy.unsafe_probability_mass(policy(generator_data.hl.states[~generator_data.hl.dones], generator_data.hl.safe_actions[~generator_data.hl.dones])).mean(), 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.ll, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c)
|
||||
if wasserstein:
|
||||
@@ -98,7 +114,8 @@ def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, v
|
||||
else:
|
||||
generator_data.ll.rewards = -F.logsigmoid(discriminator(generator_data.ll.states, generator_data.ll.actions))
|
||||
logger.add_scalar('disc/final_loss', loss, epoch)
|
||||
logger.add_scalar('disc/mean_reward_per_episode', generator_data.ll.rewards[~generator_data.ll.dones].sum() / generator_data.ll.states.shape[0], epoch)
|
||||
disc_mean_reward_per_episode = generator_data.ll.rewards[~generator_data.ll.dones].sum() / generator_data.ll.states.shape[0]
|
||||
logger.add_scalar('disc/mean_reward_per_episode', disc_mean_reward_per_episode, epoch)
|
||||
|
||||
#assert generator_data.ll.rewards.shape == generator_data.ll.dones.shape
|
||||
generator_data.hl.rewards = torch.where(~generator_data.ll.dones, generator_data.ll.rewards, torch.tensor(0.)).sum(-1)
|
||||
@@ -107,7 +124,18 @@ def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, v
|
||||
expert_data = roll_buffer(expert_data, shifts=-3, dims=0)
|
||||
|
||||
if callback is not None:
|
||||
callback(epoch, value, policy)
|
||||
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
|
||||
|
||||
@@ -119,6 +147,7 @@ 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)
|
||||
|
||||
ll_states = torch.zeros(n_episodes, max_steps_per_episode, env.max_plan_length + 1, *env.observation_space['observation'].shape)
|
||||
ll_actions = torch.zeros(n_episodes, max_steps_per_episode, env.max_plan_length + 1, *env.ll_action_space.shape)
|
||||
@@ -144,6 +173,9 @@ def rollout(env_fn, policy, n_episodes, max_steps_per_episode):
|
||||
safe_actions[:, s + 1] = torch.tensor(o['safe_actions']).clone().detach()
|
||||
rewards[:, s] = torch.tensor(r).clone().detach()
|
||||
dones[:, s + 1] = torch.tensor(d).clone().detach()
|
||||
collisions[:, s] = torch.from_numpy(np.stack([
|
||||
any(k['collision'] for k in i['ll']['infos']) for i in info
|
||||
])).detach().clone()
|
||||
|
||||
ll_states[:, s] = torch.from_numpy(np.stack([i['ll']['observations'] for i in info])).clone().detach()
|
||||
ll_actions[:, s] = torch.from_numpy(np.stack([i['ll']['actions'] for i in info])).clone().detach()
|
||||
@@ -158,7 +190,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, safe_actions, actions, rewards, dones), (ll_states, ll_actions, ll_rewards, ll_dones)
|
||||
return (states, safe_actions, actions, rewards, dones), (ll_states, ll_actions, ll_rewards, ll_dones), collisions
|
||||
|
||||
class SafeOptionsEnv(OptionsEnv):
|
||||
|
||||
@@ -176,6 +208,7 @@ class SafeOptionsEnv(OptionsEnv):
|
||||
return np.ones(len(self.options), dtype=bool)
|
||||
|
||||
plans = [self.plan(o) for o in self.options]
|
||||
plans = [np.pad(p, (0, self.max_plan_length - len(p)), constant_values=np.nan) for p in plans]
|
||||
plans = np.stack(plans)
|
||||
safe = feasible(self.env, plans, method=self.safe_actions_collision_method)
|
||||
if not safe.any():
|
||||
@@ -226,9 +259,9 @@ class SafeOptionsEnv(OptionsEnv):
|
||||
if d:
|
||||
break
|
||||
|
||||
if self.abort_unsafe_collision_method is not None and \
|
||||
not feasible(self.env, plan[k:], method=self.abort_unsafe_collision_method):
|
||||
break
|
||||
if self.abort_unsafe_collision_method is not None:
|
||||
if not feasible(self.env, plan[k:], method=self.abort_unsafe_collision_method):
|
||||
break
|
||||
|
||||
n_steps = k + 1
|
||||
return observations, actions, rewards, env_done, plan_done, infos, n_steps
|
||||
@@ -251,10 +284,10 @@ obs_max = np.array([
|
||||
[50, np.pi, 20, 20, np.pi, 1e-1],
|
||||
]).reshape(-1)
|
||||
|
||||
def NormalizedSafeOptionsEvalEnv(max_episode_steps=float('inf'), safe_actions_collision_method=None, abort_unsafe_collision_method=None, **kwargs):
|
||||
return OptionsTimeLimit(SafeOptionsEnv(Setobs(
|
||||
def NormalizedSafeOptionsEvalEnv(safe_actions_collision_method='circle', abort_unsafe_collision_method='circle', **kwargs):
|
||||
return SafeOptionsEnv(Setobs(
|
||||
TransformObservation(IntersimpleLidarFlatIncrementingAgent(
|
||||
n_rays=5,
|
||||
**kwargs,
|
||||
), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10))
|
||||
), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5), (10, 5)], safe_actions_collision_method=safe_actions_collision_method, abort_unsafe_collision_method=abort_unsafe_collision_method), max_episode_steps=max_episode_steps)
|
||||
), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5), (10, 5)], safe_actions_collision_method=safe_actions_collision_method, abort_unsafe_collision_method=abort_unsafe_collision_method)
|
||||
|
||||
@@ -9,7 +9,7 @@ class TransformObservation(gym.wrappers.TransformObservation):
|
||||
def __getattr__(self, name):
|
||||
return getattr(self.env, name)
|
||||
|
||||
class OptionsTimeLimit(gym.wrappers.TimeLimit):
|
||||
class IntersimpleTimeLimit(gym.wrappers.TimeLimit):
|
||||
def __getattr__(self, name):
|
||||
return getattr(self.env, name)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user