Compare commits
13 Commits
9c9ee8f21b
...
dev-idm-vi
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bc33b786aa | ||
|
|
dd201738cb | ||
|
|
5fb358d725 | ||
|
|
740e0ea9f4 | ||
|
|
88213e7d76 | ||
|
|
388c80007e | ||
|
|
3aaf252dbe | ||
|
|
779a0ea89f | ||
|
|
3a09a6eb7d | ||
|
|
3fa370eb8a | ||
|
|
a576f0fb18 | ||
|
|
575e299fc8 | ||
|
|
1e70303c57 |
@@ -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,
|
||||||
|
|||||||
@@ -1,7 +1,8 @@
|
|||||||
{
|
{
|
||||||
"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,
|
||||||
|
|||||||
@@ -1,7 +1,8 @@
|
|||||||
{
|
{
|
||||||
"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,
|
||||||
|
|||||||
@@ -1,7 +1,8 @@
|
|||||||
{
|
{
|
||||||
"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,
|
||||||
|
|||||||
@@ -1,7 +1,8 @@
|
|||||||
{
|
{
|
||||||
"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,
|
||||||
|
|||||||
@@ -3,7 +3,8 @@
|
|||||||
"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,
|
||||||
|
|||||||
@@ -3,7 +3,8 @@
|
|||||||
"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,
|
||||||
|
|||||||
@@ -3,7 +3,8 @@
|
|||||||
"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,
|
||||||
|
|||||||
@@ -3,7 +3,8 @@
|
|||||||
"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,
|
||||||
|
|||||||
10
cp-videos.sh
Executable file
10
cp-videos.sh
Executable file
@@ -0,0 +1,10 @@
|
|||||||
|
# cp-videos videos/ videos/icra23/
|
||||||
|
|
||||||
|
agents=( 5 27 39 43 47 53 63 81 83 87 93 96 105 113 124 127 130 134 )
|
||||||
|
|
||||||
|
for a in "${agents[@]}"
|
||||||
|
do
|
||||||
|
cp "$1/expert_agent/loc0/track0/agent${a}_ani.mp4" "$2/t${a}expert.mp4"
|
||||||
|
cp "$1/idm/loc0/track0/agent${a}_ani.mp4" "$2/t${a}idm.mp4"
|
||||||
|
cp "$1/shail/loc0/track0/agent${a}_ani.mp4" "$2/t${a}shail.mp4"
|
||||||
|
done
|
||||||
@@ -6,22 +6,24 @@ 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']:
|
||||||
env, env_kwargs ='NRasterizedRouteIncrementingAgent', {}
|
env, env_kwargs ='NRasterizedRouteIncrementingAgent', {}
|
||||||
|
elif method in ['idm']:
|
||||||
|
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
|
||||||
|
|
||||||
@@ -30,8 +32,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 +57,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])
|
||||||
|
|||||||
@@ -11,8 +11,8 @@ python -m eval_experiments --method shail --folder='test_policies/shail/expA'
|
|||||||
|
|
||||||
# Experiment B
|
# Experiment B
|
||||||
python -m eval_experiments --locations='[(0,4)]'
|
python -m eval_experiments --locations='[(0,4)]'
|
||||||
python -m eval_experiments --method idm --locations='[(0,4)]' --skip_running
|
python -m eval_experiments --method idm --locations='[(0,4)]'
|
||||||
python -m eval_experiments --method bc --folder='test_policies/bc/expB' --locations='[(0,4)]' --skip_running
|
python -m eval_experiments --method bc --folder='test_policies/bc/expB' --locations='[(0,4)]'
|
||||||
python -m eval_experiments --method gail --folder='test_policies/gail/expB' --locations='[(0,4)]' --skip_running
|
python -m eval_experiments --method gail --folder='test_policies/gail/expB' --locations='[(0,4)]'
|
||||||
python -m eval_experiments --method hail --folder='test_policies/hail/expB' --locations='[(0,4)]' --skip_running
|
python -m eval_experiments --method hail --folder='test_policies/hail/expB' --locations='[(0,4)]'
|
||||||
python -m eval_experiments --method shail --folder='test_policies/shail/expB' --locations='[(0,4)]' --skip_running
|
python -m eval_experiments --method shail --folder='test_policies/shail/expB' --locations='[(0,4)]'
|
||||||
@@ -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)
|
||||||
@@ -170,6 +172,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,
|
||||||
|
|||||||
20
generate_videos.sh
Executable file
20
generate_videos.sh
Executable 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
|
||||||
@@ -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,
|
||||||
|
|||||||
@@ -172,6 +172,7 @@ class IDMRulePolicy(BaseAlgorithm):
|
|||||||
|
|
||||||
# Update environment interaction graph with leader
|
# Update environment interaction graph with leader
|
||||||
self._env._env._graph._neighbor_dict={agent:[leader]}
|
self._env._env._graph._neighbor_dict={agent:[leader]}
|
||||||
|
self._env._update_graph = True
|
||||||
|
|
||||||
delta_v = v_ego - v[leader, 0]
|
delta_v = v_ego - v[leader, 0]
|
||||||
d_des = self.d_min + self.tau * v_ego + v_ego * delta_v / (2* (self.a_max*self.b_pref)**0.5 )
|
d_des = self.d_min + self.tau * v_ego + v_ego * delta_v / (2* (self.a_max*self.b_pref)**0.5 )
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -146,6 +151,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):
|
||||||
"""
|
"""
|
||||||
Postprocess and metrics after simulation episodes
|
Postprocess and metrics after simulation episodes
|
||||||
|
|||||||
30
src/evaluation/vec_env.py
Normal file
30
src/evaluation/vec_env.py
Normal 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)
|
||||||
@@ -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
|
||||||
@@ -85,10 +86,13 @@ 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 step(self, action, render_mode=None):
|
def render(self, mode='post'):
|
||||||
|
self.render_mode = mode
|
||||||
|
|
||||||
|
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()
|
||||||
|
|||||||
@@ -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(),
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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
|
||||||
Reference in New Issue
Block a user