adding idm override option flag, set to true. current running experiments for gail and shail experiment A to see how different times are. Since were on cpus on the cluster, guessing it will be 10x

This commit is contained in:
Arec Jamgochian
2022-08-07 16:33:09 -07:00
parent 9c9ee8f21b
commit 3a09a6eb7d
5 changed files with 21 additions and 10 deletions

View File

@@ -55,6 +55,7 @@ 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,6 +69,8 @@ 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)],[])
@@ -159,6 +162,7 @@ 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

@@ -12,16 +12,16 @@ def main(method:str='expert', folder:str=None, locations=[(0,0)], skip_running=F
policy_kwargs = {} policy_kwargs = {}
if method in ['expert', 'idm']: if method in ['expert', 'idm']:
env, env_kwargs ='NRasterizedRouteIncrementingAgent', {} env, env_kwargs ='NRasterizedRouteIncrementingAgent', {'use_idm':True}
elif method in ['bc','gail']: elif method in ['bc','gail']:
env='NormalizedContinuousEvalEnv' env='NormalizedContinuousEvalEnv'
env_kwargs={'stop_on_collision':True, 'max_episode_steps':1000} env_kwargs={'stop_on_collision':True, 'max_episode_steps':1000, 'use_idm':True}
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} env_kwargs={'stop_on_collision':True, 'max_episode_steps':1000, 'safe_actions_collision_method': None, 'abort_unsafe_collision_method': None, 'use_idm':True}
elif method in ['shail']: elif method in ['shail']:
env = 'NormalizedSafeOptionsEvalEnv' env = 'NormalizedSafeOptionsEvalEnv'
env_kwargs={'stop_on_collision':True, 'max_episode_steps':1000} env_kwargs={'stop_on_collision':True, 'max_episode_steps':1000, 'use_idm':True}
else: else:
raise NotImplementedError raise NotImplementedError

View File

@@ -53,6 +53,7 @@ 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)]
@@ -67,6 +68,7 @@ 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)
@@ -169,7 +171,8 @@ 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,

View File

@@ -58,6 +58,7 @@ 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'],
@@ -73,7 +74,9 @@ def training_function(config):
collision_penalty=0 collision_penalty=0
), ),
check_collisions=True, check_collisions=True,
stop_on_collision=config['trainenv']['stop_on_collision'], track=track, stop_on_collision=config['trainenv']['stop_on_collision'],
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'],
@@ -180,6 +183,7 @@ 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

@@ -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 --train best_configs/hail_expA.json python shail-experiment.py --test best_configs/hail_expA.json
python shail-experiment.py --train best_configs/shail_expA.json python shail-experiment.py --test 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 --train best_configs/hail_expB.json python shail-experiment.py --test best_configs/hail_expB.json
python shail-experiment.py --train best_configs/shail_expB.json python shail-experiment.py --test best_configs/shail_expB.json