Add collision and avg velocity metrics
This commit is contained in:
@@ -3,7 +3,7 @@ import pickle
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import matplotlib.pyplot as plt
|
import matplotlib.pyplot as plt
|
||||||
from torch.utils.data import DataLoader
|
from torch.utils.data import DataLoader
|
||||||
import intersim.collisions
|
from intersim import collisions
|
||||||
|
|
||||||
def metrics(filestr: str, test_dataset, policy):
|
def metrics(filestr: str, test_dataset, policy):
|
||||||
"""
|
"""
|
||||||
@@ -19,10 +19,17 @@ def metrics(filestr: str, test_dataset, policy):
|
|||||||
# b) applying the policy to observations in the test dataset
|
# b) applying the policy to observations in the test dataset
|
||||||
|
|
||||||
# load trajectory
|
# load trajectory
|
||||||
|
states = torch.load(filestr + '_sim_states.pt').detach()
|
||||||
|
lengths = torch.load(filestr + '_sim_lengths.pt').detach()
|
||||||
|
widths = torch.load(filestr + '_sim_widths.pt').detach()
|
||||||
|
xpoly = torch.load(filestr + '_sim_xpoly.pt').detach()
|
||||||
|
ypoly = torch.load(filestr + '_sim_ypoly.pt').detach()
|
||||||
|
|
||||||
# count collisions (from function in intersim.collisions)
|
# count collisions (from function in intersim.collisions)
|
||||||
|
n_collisions = collisions.count_collisions_trajectory(states, lengths, widths)
|
||||||
|
|
||||||
# calculate average velocity
|
# calculate average velocity
|
||||||
|
avg_v = average_velocity(states)
|
||||||
|
|
||||||
# calculate divergence between velocity distributions
|
# calculate divergence between velocity distributions
|
||||||
|
|
||||||
@@ -59,11 +66,18 @@ def visualize_distribution(true, pred, filestr):
|
|||||||
plt.legend(['True', 'Predicted'])
|
plt.legend(['True', 'Predicted'])
|
||||||
plt.savefig(filestr+'.png')
|
plt.savefig(filestr+'.png')
|
||||||
|
|
||||||
def average_velocity(x):
|
def average_velocity(states):
|
||||||
"""
|
"""
|
||||||
|
Compute average of average velocity over all vehicles.
|
||||||
|
Args:
|
||||||
|
states (torch.tensor): (T,nv,5) vehicle states
|
||||||
|
Returns
|
||||||
|
avg_v (float): average velocity
|
||||||
"""
|
"""
|
||||||
pass
|
velocities = states[:,:,2]
|
||||||
|
vehicle_avg_v = nanmean(velocities, dim=0)
|
||||||
|
arg_v = nanmean(vehicle_avg_v)
|
||||||
|
return arg_v
|
||||||
|
|
||||||
def divergence(p, q, type='kl'):
|
def divergence(p, q, type='kl'):
|
||||||
"""
|
"""
|
||||||
@@ -74,4 +88,23 @@ def divergence(p, q, type='kl'):
|
|||||||
Returns:
|
Returns:
|
||||||
d (float): approximate divergence
|
d (float): approximate divergence
|
||||||
"""
|
"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
def nanmean(v, *args, inplace=False, **kwargs):
|
||||||
|
"""
|
||||||
|
Calculate mean over not nan entries
|
||||||
|
|
||||||
|
To be added to torch as torch.nanmean in the next release
|
||||||
|
https://github.com/pytorch/pytorch/issues/61474, https://github.com/pytorch/pytorch/issues/21987
|
||||||
|
|
||||||
|
Args:
|
||||||
|
v (torch.tensor): arbitrary tensor
|
||||||
|
Returns:
|
||||||
|
result (torch.tensor): mean over non nan elements
|
||||||
|
"""
|
||||||
|
if not inplace:
|
||||||
|
v = v.clone()
|
||||||
|
is_nan = torch.isnan(v)
|
||||||
|
v[is_nan] = 0
|
||||||
|
result = v.sum(*args, **kwargs) / (~is_nan).float().sum(*args, **kwargs)
|
||||||
|
return result
|
||||||
Reference in New Issue
Block a user