Optionally save videos of policy evaluations

This commit is contained in:
ebuehrle
2022-03-04 15:58:18 +01:00
parent a9feec4f38
commit 1e70303c57
9 changed files with 95 additions and 15 deletions

View File

@@ -1,2 +1 @@
from src.data.expert_data import generate_expert_data, load_expert_data
from src.data.data_utils import InteractionDatasetSingleAgent

View File

@@ -177,7 +177,8 @@ def evaluate_policy(locations:List[Tuple[int,int]],
env_kwargs:dict,
method: 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.
Return metrics for that policy
@@ -230,7 +231,13 @@ def evaluate_policy(locations:List[Tuple[int,int]],
policy = load_policy(method, policy_file, policy_kwargs, eval_env)
# 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
@@ -364,7 +371,8 @@ def eval_main(
policy_kwargs: dict={},
env: str='NRasterizedRouteIncrementingAgent',
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
over all files.
@@ -410,7 +418,7 @@ def eval_main(
else:
# 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)
save_metrics(smetrics, outbase+'_summary.pkl')
cmetrics = comparison_metrics(policy_metrics, expert_metrics, outbase=outbase)

View File

@@ -9,6 +9,7 @@ from tqdm import tqdm
from src.util.wrappers import IntersimpleTimeLimit
from src.options.envs import OptionsEnv
from src.safe_options.options import SafeOptionsEnv
from src.evaluation.vec_env import CallbackWhenDoneVecEnv
class IntersimpleEvaluation:
"""
@@ -80,7 +81,7 @@ class IntersimpleEvaluation:
with open(filestr, 'wb') as 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
@@ -88,6 +89,8 @@ class IntersimpleEvaluation:
policy (BaseClass.BaseAlgorithm): policy in which policy.predict(observation)[0] returns an action
filestr (str): path-like string to dump metrics to or None
"""
self.videos_folder = videos_folder
self.reset()
if self.use_pbar:
self.pbar = tqdm(total=self.n_episodes)
@@ -97,10 +100,11 @@ class IntersimpleEvaluation:
evaluate_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,
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:
self.pbar.close()
@@ -145,6 +149,13 @@ class IntersimpleEvaluation:
if done and self.use_pbar:
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)
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.action_space = gym.spaces.Discrete(len(options))
self.max_plan_length = max(t for _, t in options)
self.render_mode = None
def plan(self, option):
target_v, t = option
@@ -84,11 +85,14 @@ class OptionsEnv(Wrapper):
n_steps = k + 1
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)
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_reward = (ll_rewards * ~ll_plan_done).sum().item()
hl_done = ll_env_done[ll_steps-1].item()

View File

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

View File

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