4 Commits

Author SHA1 Message Date
ebuehrle
3fa370eb8a Add flag to skip seeds 2022-03-05 07:12:34 +01:00
ebuehrle
a576f0fb18 Close figures 2022-03-04 16:14:35 +01:00
ebuehrle
575e299fc8 Generate videos of expert data 2022-03-04 16:07:15 +01:00
ebuehrle
1e70303c57 Optionally save videos of policy evaluations 2022-03-04 15:58:18 +01:00
90 changed files with 249 additions and 190 deletions

View File

@@ -1,22 +1,10 @@
# InteractionImitation # InteractionImitation
Imitation Learning with the [Interaction Dataset](https://interaction-dataset.com/) via the [InteractionSimulator](https://github.com/sisl/InteractionSimulator) gym environments. Imitation Learning with the INTERACTION Dataset
Code for "[SHAIL: Safety-Aware Hierarchical Adversarial Imitation Learning for Autonomous Driving in Urban Environments](https://arxiv.org/abs/2204.01922)".
If you find this repository useful, please cite the paper:
```
@article{jamgochian2022shail,
author = {Arec Jamgochian and Etienne Buehrle and Johannes Fischer and Mykel J. Kochenderfer},
title = {{SHAIL}: Safety-Aware Hierarchical Adversarial Imitation Learning for Autonomous Driving in Urban Environments},
journal = {arXiv:2204.01922 [cs]},
year = {2022}
}
```
## Getting started ## Getting started
Clone the `InteractionSimulator` with the `shail` tag and pip install the module. Clone InteractionSimulator and pip install the module.
``` ```
git clone --branch shail https://github.com/sisl/InteractionSimulator.git git clone https://github.com/sisl/InteractionSimulator.git
cd InteractionSimulator cd InteractionSimulator
pip install -e . pip install -e .
cd .. cd ..

View File

@@ -55,7 +55,6 @@ def training_function(config):
), ),
check_collisions=True, check_collisions=True,
stop_on_collision=config['trainenv']['stop_on_collision'], stop_on_collision=config['trainenv']['stop_on_collision'],
use_idm=config['trainenv']['use_idm'],
), collision_distance=6, collision_penalty=100), ), collision_distance=6, collision_penalty=100),
lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10) lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)
)) for _ in range(60)] )) for _ in range(60)]
@@ -69,8 +68,6 @@ def training_function(config):
), ),
check_collisions=True, check_collisions=True,
stop_on_collision=config['trainenv']['stop_on_collision'], stop_on_collision=config['trainenv']['stop_on_collision'],
use_idm=config['trainenv']['use_idm'],
track=track,
), collision_distance=6, collision_penalty=100), ), collision_distance=6, collision_penalty=100),
lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10) lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)
)) for _ in range(15)] for track in range(4)],[]) )) for _ in range(15)] for track in range(4)],[])
@@ -162,7 +159,6 @@ if __name__ == '__main__':
'experiment': args.train, 'experiment': args.train,
'trainenv': { 'trainenv': {
'stop_on_collision': False, 'stop_on_collision': False,
'use_idm':True,
}, },
'policy': { 'policy': {
'learning_rate': 3e-4, 'learning_rate': 3e-4,

View File

@@ -1,8 +1,7 @@
{ {
"experiment": "A", "experiment": "A",
"trainenv": { "trainenv": {
"stop_on_collision": false, "stop_on_collision": false
"use_idm": true
}, },
"policy": { "policy": {
"learning_rate": 0.0003, "learning_rate": 0.0003,

View File

@@ -1,8 +1,7 @@
{ {
"experiment": "B", "experiment": "B",
"trainenv": { "trainenv": {
"stop_on_collision": false, "stop_on_collision": false
"use_idm": true
}, },
"policy": { "policy": {
"learning_rate": 0.0003, "learning_rate": 0.0003,

View File

@@ -1,8 +1,7 @@
{ {
"experiment": "A", "experiment": "A",
"trainenv": { "trainenv": {
"stop_on_collision": false, "stop_on_collision": false
"use_idm": true
}, },
"policy": { "policy": {
"learning_rate": 0.0003, "learning_rate": 0.0003,

View File

@@ -1,8 +1,7 @@
{ {
"experiment": "B", "experiment": "B",
"trainenv": { "trainenv": {
"stop_on_collision": false, "stop_on_collision": false
"use_idm": true
}, },
"policy": { "policy": {
"learning_rate": 0.0003, "learning_rate": 0.0003,

View File

@@ -3,8 +3,7 @@
"trainenv": { "trainenv": {
"stop_on_collision": false, "stop_on_collision": false,
"safe_actions_collision_method": null, "safe_actions_collision_method": null,
"abort_unsafe_collision_method": null, "abort_unsafe_collision_method": null
"use_idm": true
}, },
"policy": { "policy": {
"learning_rate": 0.0003, "learning_rate": 0.0003,

View File

@@ -3,8 +3,7 @@
"trainenv": { "trainenv": {
"stop_on_collision": false, "stop_on_collision": false,
"safe_actions_collision_method": null, "safe_actions_collision_method": null,
"abort_unsafe_collision_method": null, "abort_unsafe_collision_method": null
"use_idm": true
}, },
"policy": { "policy": {
"learning_rate": 0.0003, "learning_rate": 0.0003,

View File

@@ -3,8 +3,7 @@
"trainenv": { "trainenv": {
"stop_on_collision": false, "stop_on_collision": false,
"safe_actions_collision_method": "circle", "safe_actions_collision_method": "circle",
"abort_unsafe_collision_method": "circle", "abort_unsafe_collision_method": "circle"
"use_idm": true
}, },
"policy": { "policy": {
"learning_rate": 0.0003, "learning_rate": 0.0003,

View File

@@ -3,8 +3,7 @@
"trainenv": { "trainenv": {
"stop_on_collision": false, "stop_on_collision": false,
"safe_actions_collision_method": "circle", "safe_actions_collision_method": "circle",
"abort_unsafe_collision_method": "circle", "abort_unsafe_collision_method": "circle"
"use_idm": true
}, },
"policy": { "policy": {
"learning_rate": 0.0003, "learning_rate": 0.0003,

View File

@@ -6,22 +6,22 @@ import json
activations = [torch.nn.Tanh, torch.nn.LeakyReLU] activations = [torch.nn.Tanh, torch.nn.LeakyReLU]
def main(method:str='expert', folder:str=None, locations=[(0,0)], skip_running=False): def main(method:str='expert', folder:str=None, locations=[(0,0)], skip_running=False, save_videos:bool=False, videos_folder:str='videos', first_seed_only:bool=False):
exclude_keys_from_policy_kwargs = {'learning_rate', 'learning_rate_decay', 'clip_ratio', 'iterations_per_epoch', 'option'} exclude_keys_from_policy_kwargs = {'learning_rate', 'learning_rate_decay', 'clip_ratio', 'iterations_per_epoch', 'option'}
policy_kwargs = {} policy_kwargs = {}
if method in ['expert', 'idm']: if method in ['expert', 'expert_agent', 'idm']:
env, env_kwargs ='NRasterizedRouteIncrementingAgent', {'use_idm':True} env, env_kwargs ='NRasterizedRouteIncrementingAgent', {}
elif method in ['bc','gail']: elif method in ['bc','gail']:
env='NormalizedContinuousEvalEnv' env='NormalizedContinuousEvalEnv'
env_kwargs={'stop_on_collision':True, 'max_episode_steps':1000, 'use_idm':True} env_kwargs={'stop_on_collision':True, 'max_episode_steps':1000}
elif method in ['hail']: elif method in ['hail']:
env = 'NormalizedSafeOptionsEvalEnv' env = 'NormalizedSafeOptionsEvalEnv'
env_kwargs={'stop_on_collision':True, 'max_episode_steps':1000, 'safe_actions_collision_method': None, 'abort_unsafe_collision_method': None, 'use_idm':True} env_kwargs={'stop_on_collision':True, 'max_episode_steps':1000, 'safe_actions_collision_method': None, 'abort_unsafe_collision_method': None}
elif method in ['shail']: elif method in ['shail']:
env = 'NormalizedSafeOptionsEvalEnv' env = 'NormalizedSafeOptionsEvalEnv'
env_kwargs={'stop_on_collision':True, 'max_episode_steps':1000, 'use_idm':True} env_kwargs={'stop_on_collision':True, 'max_episode_steps':1000}
else: else:
raise NotImplementedError raise NotImplementedError
@@ -30,8 +30,13 @@ def main(method:str='expert', folder:str=None, locations=[(0,0)], skip_running=F
if folder is not None: if folder is not None:
files = [os.path.join(folder, f) for f in os.listdir(folder) if os.path.isfile(os.path.join(folder, f))] files = [os.path.join(folder, f) for f in os.listdir(folder) if os.path.isfile(os.path.join(folder, f))]
files = [f for f in files if f.endswith('.pt')] files = [f for f in files if f.endswith('.pt')]
if first_seed_only:
files = files[:1]
with open(os.path.join(folder, 'config.json'), 'rb') as f: with open(os.path.join(folder, 'config.json'), 'rb') as f:
config = json.load(f) config = json.load(f)
print('%i policy files found in %s folder' %(len(files), folder)) print('%i policy files found in %s folder' %(len(files), folder))
print('found policy config', config['policy']) print('found policy config', config['policy'])
@@ -50,7 +55,8 @@ def main(method:str='expert', folder:str=None, locations=[(0,0)], skip_running=F
policy_file=policy_file, policy_file=policy_file,
policy_kwargs=policy_kwargs, policy_kwargs=policy_kwargs,
env=env, env=env,
env_kwargs=env_kwargs) env_kwargs=env_kwargs,
videos_folder=None if not save_videos else videos_folder)
outfolder = os.path.dirname(outbase) outfolder = os.path.dirname(outbase)
else: else:
locstr = 'loc_'+'_'.join([f'r{ro}t{tr}' for (ro,tr) in locations]) locstr = 'loc_'+'_'.join([f'r{ro}t{tr}' for (ro,tr) in locations])

View File

@@ -53,7 +53,6 @@ def training_function(config):
), ),
check_collisions=True, check_collisions=True,
stop_on_collision=config['trainenv']['stop_on_collision'], stop_on_collision=config['trainenv']['stop_on_collision'],
use_idm=config['trainenv']['use_idm'],
), collision_distance=6, collision_penalty=100), ), collision_distance=6, collision_penalty=100),
lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10) lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)
)) for _ in range(60)] )) for _ in range(60)]
@@ -68,7 +67,6 @@ def training_function(config):
), ),
check_collisions=True, check_collisions=True,
stop_on_collision=config['trainenv']['stop_on_collision'], stop_on_collision=config['trainenv']['stop_on_collision'],
use_idm=config['trainenv']['use_idm'],
track=track, track=track,
), collision_distance=6, collision_penalty=100), ), collision_distance=6, collision_penalty=100),
lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10) lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)
@@ -171,8 +169,7 @@ if __name__ == '__main__':
config={ config={
'experiment': args.train, 'experiment': args.train,
'trainenv': { 'trainenv': {
'stop_on_collision': False, 'stop_on_collision': False,
'use_idm': True,
}, },
'policy': { 'policy': {
'learning_rate': 3e-4, 'learning_rate': 3e-4,

20
generate_videos.sh Executable file
View File

@@ -0,0 +1,20 @@
# can add --skip_running if you've already run the saved policies through the test environments and have appropriate
# metrics in the out folder. Doing so will generate average metrics quickly.
# Experiment A
python -m eval_experiments
python -m eval_experiments --method expert_agent --save_videos --first_seed_only
python -m eval_experiments --method idm --save_videos --first_seed_only
python -m eval_experiments --method bc --folder='test_policies/bc/expA' --save_videos --first_seed_only
python -m eval_experiments --method gail --folder='test_policies/gail/expA' --save_videos --first_seed_only
python -m eval_experiments --method hail --folder='test_policies/hail/expA' --save_videos --first_seed_only
python -m eval_experiments --method shail --folder='test_policies/shail/expA' --save_videos --first_seed_only
# Experiment B
python -m eval_experiments --locations='[(0,4)]'
python -m eval_experiments --method expert_agent --locations='[(0,4)]' --save_videos --first_seed_only
python -m eval_experiments --method idm --locations='[(0,4)]' --save_videos --first_seed_only
python -m eval_experiments --method bc --folder='test_policies/bc/expB' --locations='[(0,4)]' --save_videos --first_seed_only
python -m eval_experiments --method gail --folder='test_policies/gail/expB' --locations='[(0,4)]' --save_videos --first_seed_only
python -m eval_experiments --method hail --folder='test_policies/hail/expB' --locations='[(0,4)]' --save_videos --first_seed_only
python -m eval_experiments --method shail --folder='test_policies/shail/expB' --locations='[(0,4)]' --save_videos --first_seed_only

View File

@@ -58,7 +58,6 @@ def training_function(config):
), ),
check_collisions=True, check_collisions=True,
stop_on_collision=config['trainenv']['stop_on_collision'], stop_on_collision=config['trainenv']['stop_on_collision'],
use_idm=config['trainenv']['use_idm'],
), collision_distance=6, collision_penalty=100), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) ), collision_distance=6, collision_penalty=100), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10))
), options=option_list[config['policy']['option']], ), options=option_list[config['policy']['option']],
safe_actions_collision_method=config['trainenv']['safe_actions_collision_method'], safe_actions_collision_method=config['trainenv']['safe_actions_collision_method'],
@@ -74,9 +73,7 @@ def training_function(config):
collision_penalty=0 collision_penalty=0
), ),
check_collisions=True, check_collisions=True,
stop_on_collision=config['trainenv']['stop_on_collision'], stop_on_collision=config['trainenv']['stop_on_collision'], track=track,
use_idm=config['trainenv']['use_idm'],
track=track,
), collision_distance=6, collision_penalty=100), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) ), collision_distance=6, collision_penalty=100), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10))
), options=option_list[config['policy']['option']], ), options=option_list[config['policy']['option']],
safe_actions_collision_method=config['trainenv']['safe_actions_collision_method'], safe_actions_collision_method=config['trainenv']['safe_actions_collision_method'],
@@ -183,7 +180,6 @@ if __name__ == '__main__':
'stop_on_collision': False, 'stop_on_collision': False,
'safe_actions_collision_method': 'circle', 'safe_actions_collision_method': 'circle',
'abort_unsafe_collision_method': 'circle', 'abort_unsafe_collision_method': 'circle',
'use_idm':True,
}, },
'policy': { 'policy': {
'learning_rate': 3e-4, 'learning_rate': 3e-4,

View File

@@ -0,0 +1 @@

View File

@@ -5,6 +5,7 @@ import intersim
from intersim.envs import Intersimple from intersim.envs import Intersimple
from stable_baselines3.common.base_class import BaseAlgorithm from stable_baselines3.common.base_class import BaseAlgorithm
from src.baselines import IDMRulePolicy from src.baselines import IDMRulePolicy
from src.data.expert import NormalizedIntersimpleExpert
from src.evaluation import IntersimpleEvaluation from src.evaluation import IntersimpleEvaluation
import src.gail.options as options_envs import src.gail.options as options_envs
from src.evaluation.metrics import divergence, visualize_distribution, rwse from src.evaluation.metrics import divergence, visualize_distribution, rwse
@@ -40,6 +41,8 @@ def load_policy(method:str,
ml = torch.device('cpu') if not torch.cuda.is_available() else None ml = torch.device('cpu') if not torch.cuda.is_available() else None
if method == 'idm': if method == 'idm':
policy = IDMRulePolicy(env, **policy_kwargs) policy = IDMRulePolicy(env, **policy_kwargs)
elif method == 'expert_agent':
policy = NormalizedIntersimpleExpert(env, **policy_kwargs)
elif method == 'bc': elif method == 'bc':
policy = SetPolicy(env.action_space.shape[-1], **policy_kwargs) policy = SetPolicy(env.action_space.shape[-1], **policy_kwargs)
policy.load_state_dict(torch.load(policy_file, map_location=ml)) policy.load_state_dict(torch.load(policy_file, map_location=ml))
@@ -177,7 +180,8 @@ def evaluate_policy(locations:List[Tuple[int,int]],
env_kwargs:dict, env_kwargs:dict,
method: str, method: str,
policy_file: str, policy_file: str,
policy_kwargs:dict) -> List[Dict[str,list]]: policy_kwargs:dict,
videos_folder: Optional[str] = None) -> List[Dict[str,list]]:
""" """
Evaluate policy on an incrementing agent environment at all locations. Evaluate policy on an incrementing agent environment at all locations.
Return metrics for that policy Return metrics for that policy
@@ -230,7 +234,13 @@ def evaluate_policy(locations:List[Tuple[int,int]],
policy = load_policy(method, policy_file, policy_kwargs, eval_env) policy = load_policy(method, policy_file, policy_kwargs, eval_env)
# run policy on environment # run policy on environment
policy_metrics[i] = evaluator.evaluate(policy) policy_videos_folder = None
if videos_folder is not None:
policy_videos_folder = os.path.join(videos_folder, method, f'loc{iround}', f'track{track}')
os.makedirs(policy_videos_folder, exist_ok=True)
policy_metrics[i] = evaluator.evaluate(
policy, videos_folder=policy_videos_folder
)
return policy_metrics return policy_metrics
@@ -364,7 +374,8 @@ def eval_main(
policy_kwargs: dict={}, policy_kwargs: dict={},
env: str='NRasterizedRouteIncrementingAgent', env: str='NRasterizedRouteIncrementingAgent',
env_kwargs: dict={}, env_kwargs: dict={},
seed: int=0): seed: int=0,
videos_folder: Optional[str]=None):
""" """
Test a particular model at different testing locations/tracks and compute average metrics Test a particular model at different testing locations/tracks and compute average metrics
over all files. over all files.
@@ -410,7 +421,7 @@ def eval_main(
else: else:
# evaluate it on the given roundabouts # evaluate it on the given roundabouts
policy_metrics = evaluate_policy(locations, env, env_kwargs, method, policy_file, policy_kwargs) policy_metrics = evaluate_policy(locations, env, env_kwargs, method, policy_file, policy_kwargs, videos_folder=videos_folder)
smetrics = summary_metrics(policy_metrics) smetrics = summary_metrics(policy_metrics)
save_metrics(smetrics, outbase+'_summary.pkl') save_metrics(smetrics, outbase+'_summary.pkl')
cmetrics = comparison_metrics(policy_metrics, expert_metrics, outbase=outbase) cmetrics = comparison_metrics(policy_metrics, expert_metrics, outbase=outbase)

View File

@@ -9,6 +9,8 @@ from tqdm import tqdm
from src.util.wrappers import IntersimpleTimeLimit from src.util.wrappers import IntersimpleTimeLimit
from src.options.envs import OptionsEnv from src.options.envs import OptionsEnv
from src.safe_options.options import SafeOptionsEnv from src.safe_options.options import SafeOptionsEnv
from src.evaluation.vec_env import CallbackWhenDoneVecEnv
import matplotlib.pyplot as plt
class IntersimpleEvaluation: class IntersimpleEvaluation:
""" """
@@ -80,7 +82,7 @@ class IntersimpleEvaluation:
with open(filestr, 'wb') as f: with open(filestr, 'wb') as f:
pickle.dump(self._metrics, f) pickle.dump(self._metrics, f)
def evaluate(self, policy, filestr: Optional[str] = None) -> Dict[str, list]: def evaluate(self, policy, filestr: Optional[str] = None, videos_folder: Optional[str] = None) -> Dict[str, list]:
""" """
Evaluate a policy on the incrementing agent evaluation environment Evaluate a policy on the incrementing agent evaluation environment
@@ -88,6 +90,8 @@ class IntersimpleEvaluation:
policy (BaseClass.BaseAlgorithm): policy in which policy.predict(observation)[0] returns an action policy (BaseClass.BaseAlgorithm): policy in which policy.predict(observation)[0] returns an action
filestr (str): path-like string to dump metrics to or None filestr (str): path-like string to dump metrics to or None
""" """
self.videos_folder = videos_folder
self.reset() self.reset()
if self.use_pbar: if self.use_pbar:
self.pbar = tqdm(total=self.n_episodes) self.pbar = tqdm(total=self.n_episodes)
@@ -97,10 +101,11 @@ class IntersimpleEvaluation:
evaluate_policy( evaluate_policy(
policy, policy,
self.env, self.env if self.videos_folder is None else CallbackWhenDoneVecEnv([lambda: self.env], self.done_callback),
n_eval_episodes=self.n_episodes, n_eval_episodes=self.n_episodes,
callback=self.evaluate_options_policy_callback if self.is_options_env else self.evaluate_policy_callback, callback=self.evaluate_options_policy_callback if self.is_options_env else self.evaluate_policy_callback,
return_episode_rewards=False return_episode_rewards=False,
render=self.videos_folder is not None,
) )
if self.use_pbar: if self.use_pbar:
self.pbar.close() self.pbar.close()
@@ -145,6 +150,14 @@ class IntersimpleEvaluation:
if done and self.use_pbar: if done and self.use_pbar:
self.pbar.update(1) self.pbar.update(1)
def done_callback(self, info):
if self.is_options_env:
info = info['ll']['infos'][0]
agent = info['agent']
filestr = os.path.join(self.videos_folder, f'agent{agent}')
self.env.close(filestr=filestr)
plt.close('all')
def post_proc(self): def post_proc(self):
""" """

30
src/evaluation/vec_env.py Normal file
View File

@@ -0,0 +1,30 @@
from stable_baselines3.common.vec_env import DummyVecEnv
from stable_baselines3.common.vec_env.base_vec_env import VecEnvStepReturn
from copy import deepcopy
import numpy as np
class CallbackWhenDoneVecEnv(DummyVecEnv):
"""DummyVecEnv that calls `done_callback` before resetting the wrapped environment."""
def __init__(self, env_fns, done_callback):
assert len(env_fns) == 1 # for now
super().__init__(env_fns)
self.done_callback = done_callback
def step_wait(self) -> VecEnvStepReturn:
for env_idx in range(self.num_envs):
obs, self.buf_rews[env_idx], self.buf_dones[env_idx], self.buf_infos[env_idx] = self.envs[env_idx].step(
self.actions[env_idx]
)
if self.buf_dones[env_idx]:
# save final observation where user can get it, then reset
self.buf_infos[env_idx]["terminal_observation"] = obs
self.done_callback(deepcopy(self.buf_infos[env_idx]))
obs = self.envs[env_idx].reset()
self._save_obs(env_idx, obs)
return (self._obs_from_buf(), np.copy(self.buf_rews), np.copy(self.buf_dones), deepcopy(self.buf_infos))
def render(self, mode='post'):
super().render(mode)

View File

@@ -45,6 +45,7 @@ class OptionsEnv(Wrapper):
self.options = options self.options = options
self.action_space = gym.spaces.Discrete(len(options)) self.action_space = gym.spaces.Discrete(len(options))
self.max_plan_length = max(t for _, t in options) self.max_plan_length = max(t for _, t in options)
self.render_mode = None
def plan(self, option): def plan(self, option):
target_v, t = option target_v, t = option
@@ -84,11 +85,14 @@ class OptionsEnv(Wrapper):
n_steps = k + 1 n_steps = k + 1
return observations, actions, rewards, env_done, plan_done, infos, n_steps return observations, actions, rewards, env_done, plan_done, infos, n_steps
def render(self, mode='post'):
self.render_mode = mode
def step(self, action, render_mode=None): def step(self, action):
a = int(action) a = int(action)
assert a == 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) 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], self.render_mode)
hl_obs = ll_obs[ll_steps] hl_obs = ll_obs[ll_steps]
hl_reward = (ll_rewards * ~ll_plan_done).sum().item() hl_reward = (ll_rewards * ~ll_plan_done).sum().item()
hl_done = ll_env_done[ll_steps-1].item() hl_done = ll_env_done[ll_steps-1].item()

View File

@@ -225,8 +225,8 @@ class SafeOptionsEnv(OptionsEnv):
} }
return obs return obs
def step(self, action, render_mode=None): def step(self, action):
obs, reward, done, info = super().step(action, render_mode) obs, reward, done, info = super().step(action)
obs = { obs = {
'observation': obs, 'observation': obs,
'safe_actions': self.safe_actions(), 'safe_actions': self.safe_actions(),

View File

@@ -5,14 +5,23 @@ class Wrapper(gym.Wrapper):
def __getattr__(self, name): def __getattr__(self, name):
return getattr(self.env, name) return getattr(self.env, name)
def close(self, *args, **kwargs):
return self.env.close(*args, **kwargs)
class TransformObservation(gym.wrappers.TransformObservation): class TransformObservation(gym.wrappers.TransformObservation):
def __getattr__(self, name): def __getattr__(self, name):
return getattr(self.env, name) return getattr(self.env, name)
def close(self, *args, **kwargs):
return self.env.close(*args, **kwargs)
class IntersimpleTimeLimit(gym.wrappers.TimeLimit): class IntersimpleTimeLimit(gym.wrappers.TimeLimit):
def __getattr__(self, name): def __getattr__(self, name):
return getattr(self.env, name) return getattr(self.env, name)
def close(self, *args, **kwargs):
return self.env.close(*args, **kwargs)
class CollisionPenaltyWrapper(Wrapper): class CollisionPenaltyWrapper(Wrapper):
def __init__(self, env, collision_distance, collision_penalty, *args, **kwargs): def __init__(self, env, collision_distance, collision_penalty, *args, **kwargs):

View File

@@ -1,32 +1,31 @@
{ {
"discriminator": { "experiment": "A",
"activation": 0, "trainenv": {
"hidden_layer_size": 10, "stop_on_collision": false
"iterations_per_epoch": 100, },
"learning_rate": 0.001, "policy": {
"n_hidden_layers_element": 4, "learning_rate": 0.0003,
"n_hidden_layers_global": 1, "learning_rate_decay": 1.0,
"weight_decay": 0.0001 "clip_ratio": 0.2,
}, "iterations_per_epoch": 100,
"experiment": "A", "hidden_layer_size": 40,
"policy": { "n_hidden_layers": 2,
"activation": 0, "activation": 0
"clip_ratio": 0.2, },
"hidden_layer_size": 40, "value": {
"iterations_per_epoch": 100, "learning_rate": 0.0001,
"learning_rate": 0.0003, "weight_decay": 0.001,
"learning_rate_decay": 1.0, "iterations_per_epoch": 1000
"n_hidden_layers": 2 },
}, "discriminator": {
"seed": 5, "learning_rate": 0.001,
"train_epochs": 100, "weight_decay": 0.0001,
"trainenv": { "iterations_per_epoch": 100,
"stop_on_collision": false, "n_hidden_layers_element": 4,
"use_idm": true "n_hidden_layers_global": 1,
}, "hidden_layer_size": 10,
"value": { "activation": 0
"iterations_per_epoch": 1000, },
"learning_rate": 0.0001, "train_epochs": 100,
"weight_decay": 0.001 "seed": 0
}
} }

View File

@@ -1,32 +1,31 @@
{ {
"discriminator": { "experiment": "B",
"activation": 0, "trainenv": {
"hidden_layer_size": 10, "stop_on_collision": false
"iterations_per_epoch": 100, },
"learning_rate": 0.001, "policy": {
"n_hidden_layers_element": 4, "learning_rate": 0.0003,
"n_hidden_layers_global": 1, "learning_rate_decay": 1.0,
"weight_decay": 0.0001 "clip_ratio": 0.2,
}, "iterations_per_epoch": 100,
"experiment": "B", "hidden_layer_size": 40,
"policy": { "n_hidden_layers": 2,
"activation": 0, "activation": 0
"clip_ratio": 0.2, },
"hidden_layer_size": 40, "value": {
"iterations_per_epoch": 100, "learning_rate": 0.0001,
"learning_rate": 0.0003, "weight_decay": 0.001,
"learning_rate_decay": 1.0, "iterations_per_epoch": 1000
"n_hidden_layers": 2 },
}, "discriminator": {
"seed": 4, "learning_rate": 0.001,
"train_epochs": 100, "weight_decay": 0.0001,
"trainenv": { "iterations_per_epoch": 100,
"stop_on_collision": false, "n_hidden_layers_element": 4,
"use_idm": true "n_hidden_layers_global": 1,
}, "hidden_layer_size": 10,
"value": { "activation": 0
"iterations_per_epoch": 1000, },
"learning_rate": 0.0001, "train_epochs": 100,
"weight_decay": 0.001 "seed": 0
}
} }

View File

@@ -1,34 +1,33 @@
{ {
"discriminator": { "experiment": "A",
"activation": 0, "trainenv": {
"hidden_layer_size": 10, "stop_on_collision": false,
"iterations_per_epoch": 100, "safe_actions_collision_method": "circle",
"learning_rate": 0.001, "abort_unsafe_collision_method": "circle"
"n_hidden_layers_element": 4, },
"n_hidden_layers_global": 1, "policy": {
"weight_decay": 0.0001 "learning_rate": 0.0003,
}, "learning_rate_decay": 1.0,
"experiment": "A", "clip_ratio": 0.2,
"policy": { "iterations_per_epoch": 100,
"activation": 0, "hidden_layer_size": 40,
"clip_ratio": 0.2, "n_hidden_layers": 2,
"hidden_layer_size": 40, "activation": 0,
"iterations_per_epoch": 100, "option": 0
"learning_rate": 0.0003, },
"learning_rate_decay": 1.0, "value": {
"n_hidden_layers": 2, "learning_rate": 0.001,
"option": 0 "iterations_per_epoch": 1000
}, },
"seed": 3, "discriminator": {
"train_epochs": 90, "learning_rate": 0.001,
"trainenv": { "weight_decay": 0.0001,
"abort_unsafe_collision_method": "circle", "iterations_per_epoch": 100,
"safe_actions_collision_method": "circle", "n_hidden_layers_element": 4,
"stop_on_collision": false, "n_hidden_layers_global": 1,
"use_idm": true "hidden_layer_size": 10,
}, "activation": 0
"value": { },
"iterations_per_epoch": 1000, "train_epochs": 90,
"learning_rate": 0.001 "seed": 0
}
} }

View File

@@ -1,34 +1,33 @@
{ {
"discriminator": { "experiment": "B",
"activation": 0, "trainenv": {
"hidden_layer_size": 10, "stop_on_collision": false,
"iterations_per_epoch": 100, "safe_actions_collision_method": "circle",
"learning_rate": 0.001, "abort_unsafe_collision_method": "circle"
"n_hidden_layers_element": 4, },
"n_hidden_layers_global": 2, "policy": {
"weight_decay": 0.0001 "learning_rate": 0.0003,
}, "learning_rate_decay": 1.0,
"experiment": "B", "clip_ratio": 0.2,
"policy": { "iterations_per_epoch": 100,
"activation": 0, "hidden_layer_size": 20,
"clip_ratio": 0.2, "n_hidden_layers": 2,
"hidden_layer_size": 20, "activation": 0,
"iterations_per_epoch": 100, "option": 0
"learning_rate": 0.0003, },
"learning_rate_decay": 1.0, "value": {
"n_hidden_layers": 2, "learning_rate": 0.001,
"option": 0 "iterations_per_epoch": 1000
}, },
"seed": 3, "discriminator": {
"train_epochs": 85, "learning_rate": 0.001,
"trainenv": { "weight_decay": 0.0001,
"abort_unsafe_collision_method": "circle", "iterations_per_epoch": 100,
"safe_actions_collision_method": "circle", "n_hidden_layers_element": 4,
"stop_on_collision": false, "n_hidden_layers_global": 2,
"use_idm": true "hidden_layer_size": 10,
}, "activation": 0
"value": { },
"iterations_per_epoch": 1000, "train_epochs": 85,
"learning_rate": 0.001 "seed": 0
}
} }

View File

@@ -18,11 +18,11 @@ python shail-experiment.py --train B
# Experiment A # Experiment A
python bc-experiment.py --test best_configs/bc_expA.json python bc-experiment.py --test best_configs/bc_expA.json
python gail-experiment.py --test best_configs/gail_expA.json python gail-experiment.py --test best_configs/gail_expA.json
python shail-experiment.py --test best_configs/hail_expA.json python shail-experiment.py --train best_configs/hail_expA.json
python shail-experiment.py --test best_configs/shail_expA.json python shail-experiment.py --train best_configs/shail_expA.json
# Experiment B # Experiment B
python bc-experiment.py --test best_configs/bc_expB.json python bc-experiment.py --test best_configs/bc_expB.json
python gail-experiment.py --test best_configs/gail_expB.json python gail-experiment.py --test best_configs/gail_expB.json
python shail-experiment.py --test best_configs/hail_expB.json python shail-experiment.py --train best_configs/hail_expB.json
python shail-experiment.py --test best_configs/shail_expB.json python shail-experiment.py --train best_configs/shail_expB.json