diff --git a/src/evaluation/evaluation.py b/src/evaluation/evaluation.py index cbbd0da..9e5423f 100644 --- a/src/evaluation/evaluation.py +++ b/src/evaluation/evaluation.py @@ -5,14 +5,19 @@ from stable_baselines3.common.evaluation import evaluate_policy from intersim.envs.intersimple import Intersimple from src.evaluation.metrics import nanmean, divergence, visualize_distribution +import os + class Evaluation: - def __init__(self, eval_env, n_eval_episodes=10): + def __init__(self, filestr, eval_env, expert_data, n_eval_episodes=10): # if env is a VecEnv, the code needs to be adapted, since the callback will be called after each step, # so transitions of different envs will be mixed and the total number of episodes could be larger than n_eval_episodes! assert not isinstance(eval_env, VecEnv) + self.filestr = filestr self.env = eval_env self.n_eval_episodes = n_eval_episodes + self.expert_data = expert_data + self.compute_expert_features(expert_data) self.reset() def reset(self): @@ -21,7 +26,16 @@ class Evaluation: self._episode_done = True self._accelerations = [] - def evaluate(self, epoch, generator, discriminator, expert_data): + def compute_expert_features(self, expert_data): + # expert velocities + extract_state = lambda info: info['projected_state'][info['agent']] + expert_velocities = torch.stack([extract_state(info) for info in expert_data.infos])[:,2] + self.expert_velocities = expert_velocities[~torch.isnan(expert_velocities)] + # expert accelerations + extract_accel = lambda info: info['action_taken'][info['agent']] + self.expert_accelerations = torch.cat([extract_accel(info) for info in expert_data.infos]) + + def evaluate(self, epoch, generator, discriminator): self.reset() metrics = {} @@ -43,23 +57,15 @@ class Evaluation: # if episodes terminate without collisions, then the state is fully nan policy_velocities = policy_velocities[~torch.isnan(policy_velocities)] - # expert velocities - extract_state = lambda info: info['projected_state'][info['agent']] - expert_velocities = torch.stack([extract_state(info) for info in expert_data.infos])[:,2] - expert_velocities = expert_velocities[~torch.isnan(expert_velocities)] - - metrics['avg_velocity_loss'] = (expert_velocities.mean() - policy_velocities.mean()).item() - metrics['velocity_divergence'] = divergence(policy_velocities, expert_velocities, type='js') + metrics['avg_velocity_loss'] = (self.expert_velocities.mean() - policy_velocities.mean()).item() + metrics['velocity_divergence'] = divergence(policy_velocities, self.expert_velocities, type='js') # accelerations produced by generator policy_accelerations = torch.tensor(self._accelerations) - # expert accelerations - extract_accel = lambda info: info['action_taken'][info['agent']] - expert_accelerations = torch.cat([extract_accel(info) for info in expert_data.infos]) - metrics['acceleration_divergence'] = divergence(policy_accelerations, expert_accelerations, type='js') - visualize_distribution(expert_accelerations, policy_accelerations, 'output/_action_viz{:02}'.format(epoch)) + metrics['acceleration_divergence'] = divergence(policy_accelerations, self.expert_accelerations, type='js') + visualize_distribution(self.expert_accelerations, policy_accelerations, os.path.join(self.filestr, '_action_viz{:02}'.format(epoch)) print(metrics) return metrics