precompute expert features

This commit is contained in:
Johannes Fischer
2021-10-28 18:16:00 +02:00
parent 070b8fc785
commit 0077c24074

View File

@@ -5,14 +5,19 @@ from stable_baselines3.common.evaluation import evaluate_policy
from intersim.envs.intersimple import Intersimple from intersim.envs.intersimple import Intersimple
from src.evaluation.metrics import nanmean, divergence, visualize_distribution from src.evaluation.metrics import nanmean, divergence, visualize_distribution
import os
class Evaluation: 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, # 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! # 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) assert not isinstance(eval_env, VecEnv)
self.filestr = filestr
self.env = eval_env self.env = eval_env
self.n_eval_episodes = n_eval_episodes self.n_eval_episodes = n_eval_episodes
self.expert_data = expert_data
self.compute_expert_features(expert_data)
self.reset() self.reset()
def reset(self): def reset(self):
@@ -21,7 +26,16 @@ class Evaluation:
self._episode_done = True self._episode_done = True
self._accelerations = [] 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() self.reset()
metrics = {} metrics = {}
@@ -43,23 +57,15 @@ class Evaluation:
# if episodes terminate without collisions, then the state is fully nan # if episodes terminate without collisions, then the state is fully nan
policy_velocities = policy_velocities[~torch.isnan(policy_velocities)] policy_velocities = policy_velocities[~torch.isnan(policy_velocities)]
# expert velocities metrics['avg_velocity_loss'] = (self.expert_velocities.mean() - policy_velocities.mean()).item()
extract_state = lambda info: info['projected_state'][info['agent']] metrics['velocity_divergence'] = divergence(policy_velocities, self.expert_velocities, type='js')
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')
# accelerations produced by generator # accelerations produced by generator
policy_accelerations = torch.tensor(self._accelerations) 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') metrics['acceleration_divergence'] = divergence(policy_accelerations, self.expert_accelerations, type='js')
visualize_distribution(expert_accelerations, policy_accelerations, 'output/_action_viz{:02}'.format(epoch)) visualize_distribution(self.expert_accelerations, policy_accelerations, os.path.join(self.filestr, '_action_viz{:02}'.format(epoch))
print(metrics) print(metrics)
return metrics return metrics