removing gail-trpo since performance is about the same as gail, adding experiment evaluation script, updating metric averaging to work
This commit is contained in:
74
eval_experiments.py
Normal file
74
eval_experiments.py
Normal file
@@ -0,0 +1,74 @@
|
||||
import os
|
||||
from src.eval_main import eval_main
|
||||
from src.evaluation.utils import load_and_average
|
||||
|
||||
def main(method:str='expert', folder:str=None, locations=[(0,0)], skip_running=False):
|
||||
|
||||
policy_kwargs = {}
|
||||
if method in ['expert', 'idm']:
|
||||
env, env_kwargs ='NRasterizedRouteIncrementingAgent', {}
|
||||
elif method in ['bc','gail']:
|
||||
env='NormalizedContinuousEvalEnv'
|
||||
env_kwargs={stop_on_collision:True, max_episode_steps:1000}
|
||||
elif method in ['hail']:
|
||||
env = 'NormalizedOptionsEvalEnv'
|
||||
env_kwargs={stop_on_collision:True, max_episode_steps:1000}
|
||||
elif method in ['shail']:
|
||||
env = 'NormalizedSafeOptionsEvalEnv'
|
||||
env_kwargs={stop_on_collision:True, max_episode_steps:1000}
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
files = ['']
|
||||
|
||||
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))]
|
||||
print('%i folders found in %s folder' %(len(files), folder))
|
||||
|
||||
if not skip_running:
|
||||
for policy_file in files:
|
||||
# run metrics on that file
|
||||
outbase = eval_main(locations=locations,
|
||||
method=method,
|
||||
policy_file=policy_file,
|
||||
policy_kwargs=policy_kwargs,
|
||||
env=env,
|
||||
env_kwargs=env_kwargs)
|
||||
outfolder = os.path.dirname(outbase)
|
||||
else:
|
||||
locstr = 'loc_'+'_'.join([f'r{ro}t{tr}' for (ro,tr) in locations])
|
||||
outfolder = os.path.join('out',method,locstr)
|
||||
|
||||
import pdb
|
||||
pdb.set_trace()
|
||||
# load metrics from save_path
|
||||
average_metrics = load_and_average(outfolder)
|
||||
if method in ['expert', 'idm']:
|
||||
latex_print(average_metrics, light=True)
|
||||
else:
|
||||
latex_print(average_metrics)
|
||||
|
||||
def latex_print(am, light=False):
|
||||
"""
|
||||
print latex line
|
||||
|
||||
am (Dict[str,tuple]): dict mapping metric_name to (mean, std)
|
||||
"""
|
||||
|
||||
print('success rate, distance travelled, RWSE_10, |DeltaV|, AccelJSD')
|
||||
if light:
|
||||
print("%2.1f& %2.1f & --- & --- & "
|
||||
"--- \\\\" %( 100*am['success rate'][0], am['mean travel distance'][0]))
|
||||
return
|
||||
|
||||
print("%2.1f \\scriptstyle\\pm %2.1f & %2.1f \\scriptstyle\\pm %2.1f & "
|
||||
"%1.2f \\scriptstyle\\pm %1.2f & %2.1f \\scriptstyle\\pm %1.1f & "
|
||||
"%0.3f \\scriptstyle\\pm %0.3f \\\\" %( 100*am['success rate'][0], 100*am['success rate'][1],
|
||||
am['mean travel distance'][0] , am['mean travel distance'][1] ,
|
||||
am['rwse_10s'][0] , am['rwse_10s'][1] ,
|
||||
am['average absolute average velocity'][0] , am['average absolute average velocity'][1] ,
|
||||
am['acceleration distribution divergence'][0] , am['acceleration distribution divergence'][1] ))
|
||||
|
||||
if __name__=='__main__':
|
||||
import fire
|
||||
fire.Fire(main)
|
||||
Reference in New Issue
Block a user