precompute expert features
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user