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 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
|
||||||
|
|||||||
Reference in New Issue
Block a user