From 575e299fc8dc0742d38c2b4061d72d312547f07a Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Fri, 4 Mar 2022 16:07:15 +0100 Subject: [PATCH] Generate videos of expert data --- eval_experiments.py | 2 +- generate_videos.sh | 2 ++ src/eval_main.py | 3 +++ 3 files changed, 6 insertions(+), 1 deletion(-) diff --git a/eval_experiments.py b/eval_experiments.py index bb552fd..a30fbe0 100644 --- a/eval_experiments.py +++ b/eval_experiments.py @@ -11,7 +11,7 @@ def main(method:str='expert', folder:str=None, locations=[(0,0)], skip_running=F exclude_keys_from_policy_kwargs = {'learning_rate', 'learning_rate_decay', 'clip_ratio', 'iterations_per_epoch', 'option'} policy_kwargs = {} - if method in ['expert', 'idm']: + if method in ['expert', 'expert_agent', 'idm']: env, env_kwargs ='NRasterizedRouteIncrementingAgent', {} elif method in ['bc','gail']: env='NormalizedContinuousEvalEnv' diff --git a/generate_videos.sh b/generate_videos.sh index d3b94d5..49f1103 100755 --- a/generate_videos.sh +++ b/generate_videos.sh @@ -3,6 +3,7 @@ # Experiment A python -m eval_experiments +python -m eval_experiments --method expert_agent --save_videos python -m eval_experiments --method idm --save_videos python -m eval_experiments --method bc --folder='test_policies/bc/expA' --save_videos python -m eval_experiments --method gail --folder='test_policies/gail/expA' --save_videos @@ -11,6 +12,7 @@ python -m eval_experiments --method shail --folder='test_policies/shail/expA' -- # Experiment B python -m eval_experiments --locations='[(0,4)]' +python -m eval_experiments --method expert_agent --locations='[(0,4)]' --save_videos python -m eval_experiments --method idm --locations='[(0,4)]' --save_videos python -m eval_experiments --method bc --folder='test_policies/bc/expB' --locations='[(0,4)]' --save_videos python -m eval_experiments --method gail --folder='test_policies/gail/expB' --locations='[(0,4)]' --save_videos diff --git a/src/eval_main.py b/src/eval_main.py index d60f3fa..ba6c251 100644 --- a/src/eval_main.py +++ b/src/eval_main.py @@ -5,6 +5,7 @@ import intersim from intersim.envs import Intersimple from stable_baselines3.common.base_class import BaseAlgorithm from src.baselines import IDMRulePolicy +from src.data.expert import NormalizedIntersimpleExpert from src.evaluation import IntersimpleEvaluation import src.gail.options as options_envs 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 if method == 'idm': policy = IDMRulePolicy(env, **policy_kwargs) + elif method == 'expert_agent': + policy = NormalizedIntersimpleExpert(env, **policy_kwargs) elif method == 'bc': policy = SetPolicy(env.action_space.shape[-1], **policy_kwargs) policy.load_state_dict(torch.load(policy_file, map_location=ml))