Add SHAIL

This commit is contained in:
ebuehrle
2022-02-17 23:49:06 +01:00
parent 9de6bfe9a3
commit 1624e1a349
11 changed files with 79 additions and 98 deletions

Binary file not shown.

View File

@@ -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}'

View File

@@ -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:

View File

@@ -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):

View File

@@ -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):

View File

@@ -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

View File

@@ -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']

View File

@@ -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)

View File

@@ -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]

View File

@@ -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):

View File

@@ -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):