diff --git a/checkpoints/sgail-options-setobs2.pt b/checkpoints/sgail-options-setobs2.pt new file mode 100644 index 0000000..5bc3e99 Binary files /dev/null and b/checkpoints/sgail-options-setobs2.pt differ diff --git a/evaluate_models.sh b/evaluate_models.sh index a2d8952..c4e5f72 100755 --- a/evaluate_models.sh +++ b/evaluate_models.sh @@ -13,6 +13,9 @@ python -m src.eval_main # idm python -m src.eval_main --method=idm +# behavior cloning +python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' + # GAIL python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' @@ -22,5 +25,5 @@ python -m src.eval_main --method=ogail --policy_file='checkpoints/gail-options-s # options GAIL-PPO python -m src.eval_main --method=ogail-ppo --policy_file='checkpoints/gail-ppo-options-setobs2-Feb15_22-05-38.pt' --env='NormalizedOptionsEvalEnv' --env_kwargs='{stop_on_collision:True}' -# behavior cloning -python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' +# SHAIL +python -m src.eval_main --method=sgail --policy_file='checkpoints/sgail-options-setobs2.pt' --env='NormalizedSafeOptionsEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' diff --git a/src/core/gail.py b/src/core/gail.py index d630bd4..2e34333 100644 --- a/src/core/gail.py +++ b/src/core/gail.py @@ -1,10 +1,10 @@ import torch import torch.nn.functional as F from dataclasses import dataclass -from core.reparam_module import ReparamPolicy -from core.sampling import rollout -from core.trpo import trpo_step -from core.ppo import ppo_step +from src.core.reparam_module import ReparamPolicy +from src.core.sampling import rollout +from src.core.trpo import trpo_step +from src.core.ppo import ppo_step from tqdm import tqdm class TerminalLogger: diff --git a/src/core/ppo.py b/src/core/ppo.py index c742061..c73db2b 100644 --- a/src/core/ppo.py +++ b/src/core/ppo.py @@ -1,6 +1,6 @@ import torch -from core.sampling import rollout -from core.value_estimation import gae +from src.core.sampling import rollout +from src.core.value_estimation import gae def ppo(env_fn, value, policy, epochs, rollout_episodes, rollout_steps, gamma, gae_lambda, clip_ratio, pi_opt, pi_iters, v_opt, v_iters, target_kl=None, max_grad_norm=None): diff --git a/src/core/trpo.py b/src/core/trpo.py index d91e363..67d0113 100644 --- a/src/core/trpo.py +++ b/src/core/trpo.py @@ -1,8 +1,8 @@ import torch -from core.reparam_module import ReparamPolicy -from core.sampling import rollout -from core.value_estimation import gae -from core.optimization import conjugate_gradient, line_search +from src.core.reparam_module import ReparamPolicy +from src.core.sampling import rollout +from src.core.value_estimation import gae +from src.core.optimization import conjugate_gradient, line_search def trpo(env_fn, value, policy, epochs, rollout_episodes, rollout_steps, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1): diff --git a/src/eval_main.py b/src/eval_main.py index 6b460b3..381b061 100644 --- a/src/eval_main.py +++ b/src/eval_main.py @@ -11,6 +11,8 @@ from src.evaluation.metrics import divergence, visualize_distribution from src.core.policy import SetPolicy, SetDiscretePolicy from src.core.reparam_module import ReparamPolicy 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 typing import Optional, List, Dict, Tuple import torch @@ -62,8 +64,14 @@ def load_policy(method:str, policy.load_state_dict(torch.load(policy_file)) policy.eval() elif method == 'sgail': - policy = sb3.PPO.load(policy_file) - raise NotImplementedError + policy = SetMaskedDiscretePolicy(env.action_space.n) + policy( + torch.zeros(env.observation_space['observation'].shape), + torch.zeros(env.observation_space['safe_actions'].shape) + ) + policy = ReparamPolicy(policy) + policy.load_state_dict(torch.load(policy_file)) + policy.eval() else: raise NotImplementedError return policy @@ -181,6 +189,7 @@ def evaluate_policy(locations:List[Tuple[int,int]], envs_dict = dict(intersim.envs.intersimple.__dict__) envs_dict.update(dict(options_envs.__dict__)) envs_dict.update(dict(options_envs2.__dict__)) + envs_dict.update(dict(options_envs3.__dict__)) policy_metrics = [None]* len(locations) # iterate through vehicles diff --git a/src/evaluation/evaluation.py b/src/evaluation/evaluation.py index 27817ff..96b8437 100644 --- a/src/evaluation/evaluation.py +++ b/src/evaluation/evaluation.py @@ -7,6 +7,7 @@ import os import pickle from tqdm import tqdm from src.options.envs import OptionsEnv +from src.util.wrappers import OptionsTimeLimit class IntersimpleEvaluation: """ @@ -35,7 +36,7 @@ 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) + self.is_options_env = isinstance(self.env, (OptionsEnv, OptionsTimeLimit)) # metrics present on every step of every episode self.metric_keys_all = ['v_all', 'a_all', 'col_all'] diff --git a/src/safe_options/options.py b/src/safe_options/options.py index a8d0e50..eeb5bae 100644 --- a/src/safe_options/options.py +++ b/src/safe_options/options.py @@ -3,14 +3,18 @@ import numpy as np import torch from stable_baselines3.common.vec_env import DummyVecEnv as VecEnv -from core.reparam_module import ReparamPolicy +from src.core.reparam_module import ReparamPolicy from tqdm import tqdm -from core.gail import train_discriminator, roll_buffer, TerminalLogger +from src.core.gail import train_discriminator, roll_buffer, TerminalLogger from dataclasses import dataclass -from safe_options.policy_gradient import trpo_step, ppo_step +from src.safe_options.policy_gradient import trpo_step, ppo_step import torch.nn.functional as F -from safe_options.collisions import feasible +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 @dataclass class Buffer: @@ -156,81 +160,6 @@ def rollout(env_fn, policy, n_episodes, max_steps_per_episode): return (states, safe_actions, actions, rewards, dones), (ll_states, ll_actions, ll_rewards, ll_dones) -class OptionsEnv(gym.Wrapper): - - def __init__(self, env, options): - super().__init__(env) - self.ll_action_space = env.action_space - self.options = options - self.action_space = gym.spaces.Discrete(len(options)) - self.max_plan_length = max(t for _, t in options) - - def plan(self, option): - target_v, t = option - current_v = self.env._env.state[self.env._agent, 1].item() - dt = self.env._env._dt - a = (target_v - current_v) / (t * dt) - a = self.env._normalize(a) - a = a * np.ones((t,)) - a += 0.01 * np.random.randn(*a.shape) - a = np.clip(a, self.ll_action_space.low, self.ll_action_space.high) - return a - - def execute_plan(self, obs, option, render_mode=None): - observations = np.zeros((self.max_plan_length + 1, *self.env.observation_space.shape)) - actions = np.zeros((self.max_plan_length + 1, *self.ll_action_space.shape)) - rewards = np.zeros((self.max_plan_length + 1,)) - env_done = np.ones((self.max_plan_length + 1,), dtype=bool) - plan_done = np.ones((self.max_plan_length + 1,), dtype=bool) - infos = [] - - plan = self.plan(option) - observations[0] = obs - env_done[0] = False - for k, u in enumerate(plan): - plan_done[k] = False - o, r, d, i = self.env.step(u) - actions[k] = u - rewards[k] = r - env_done[k+1] = d - infos.append(i) - observations[k+1] = o - - if render_mode is not None: - self.env.render(render_mode) - - if d: - break - - n_steps = k + 1 - return observations, actions, rewards, env_done, plan_done, infos, n_steps - - def step(self, action, render_mode=None): - a = int(action) - assert a == action - ll_obs, ll_actions, ll_rewards, ll_env_done, ll_plan_done, ll_infos, ll_steps = self.execute_plan(self.last_obs, self.options[a], render_mode) - hl_obs = ll_obs[ll_steps] - hl_reward = (ll_rewards * ~ll_plan_done).sum().item() - hl_done = ll_env_done[ll_steps].item() - hl_infos = { - 'll': { - 'observations': ll_obs, - 'actions': ll_actions, - 'rewards': ll_rewards, - 'env_done': ll_env_done, - 'plan_done': ll_plan_done, - 'infos': ll_infos, - 'steps': ll_steps, - } - } - self.last_obs = hl_obs - return hl_obs, hl_reward, hl_done, hl_infos - - def reset(self, *args, **kwargs): - self.last_obs = super().reset(*args, **kwargs) - return self.last_obs - - class SafeOptionsEnv(OptionsEnv): def __init__(self, env, options, safe_actions_collision_method=None, abort_unsafe_collision_method=None): @@ -287,7 +216,7 @@ class SafeOptionsEnv(OptionsEnv): o, r, d, i = self.env.step(u) actions[k] = u rewards[k] = r - env_done[k+1] = d + env_done[k] = d infos.append(i) observations[k+1] = o @@ -303,3 +232,29 @@ class SafeOptionsEnv(OptionsEnv): n_steps = k + 1 return observations, actions, rewards, env_done, plan_done, infos, n_steps + +obs_min = np.array([ + [-1000, -1000, 0, -np.pi, -1e-1, 0.], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], +]).reshape(-1) + +obs_max = np.array([ + [1000, 1000, 20, np.pi, 1e-1, 0.], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [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( + 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)], safe_actions_collision_method=safe_actions_collision_method, abort_unsafe_collision_method=abort_unsafe_collision_method), max_episode_steps=max_episode_steps) diff --git a/src/safe_options/policy.py b/src/safe_options/policy.py index 4e0e9d0..0325f59 100644 --- a/src/safe_options/policy.py +++ b/src/safe_options/policy.py @@ -2,7 +2,7 @@ import torch import torch.nn as nn from torch.distributions import Categorical from torch.distributions.kl import kl_divergence -from core.policy import SetDiscretePolicy +from src.core.policy import SetDiscretePolicy class SetMaskedDiscretePolicy(SetDiscretePolicy): @@ -21,6 +21,15 @@ class SetMaskedDiscretePolicy(SetDiscretePolicy): a = super().torch_dist(logits).probs return (a * (1 - z)).sum(-1) + def predict(self, observations, state=None, episode_start=None, deterministic=True): + observation = torch.tensor(observations['observation']) + safe_actions = torch.tensor(observations['safe_actions']) + if deterministic: + _, actions = self.forward(observation, safe_actions).max(-1) + else: + actions = self.sample(self.forward(observation, safe_actions)) + return actions, None + # def torch_dist_nomask(self, dist): # print('no mask logprob') # logits = dist[..., :self.action_dim] diff --git a/src/safe_options/policy_gradient.py b/src/safe_options/policy_gradient.py index 6e968e4..803af15 100644 --- a/src/safe_options/policy_gradient.py +++ b/src/safe_options/policy_gradient.py @@ -1,6 +1,6 @@ import torch -from core.value_estimation import gae -from core.optimization import conjugate_gradient, line_search +from src.core.value_estimation import gae +from src.core.optimization import conjugate_gradient, line_search def trpo_step(value, policy, states, safe_actions, actions, rewards, dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1): diff --git a/src/util/wrappers.py b/src/util/wrappers.py index 3916e75..3a088f9 100644 --- a/src/util/wrappers.py +++ b/src/util/wrappers.py @@ -9,6 +9,10 @@ class TransformObservation(gym.wrappers.TransformObservation): def __getattr__(self, name): return getattr(self.env, name) +class OptionsTimeLimit(gym.wrappers.TimeLimit): + def __getattr__(self, name): + return getattr(self.env, name) + class CollisionPenaltyWrapper(Wrapper): def __init__(self, env, collision_distance, collision_penalty, *args, **kwargs):