From 2da0e05782ebb3447ef1763c8eb41d215409fa47 Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Mon, 21 Feb 2022 10:40:52 +0100 Subject: [PATCH] Implement rwse --- src/evaluation/metrics.py | 19 +++++++++++++++++-- 1 file changed, 17 insertions(+), 2 deletions(-) diff --git a/src/evaluation/metrics.py b/src/evaluation/metrics.py index 477a22b..4098229 100644 --- a/src/evaluation/metrics.py +++ b/src/evaluation/metrics.py @@ -9,7 +9,7 @@ from typing import List def rwse(expert:List[np.ndarray], policy:List[np.ndarray], dt:float=0.1) -> float: """ - Calculated time-weighted + Calculate average mean squared displacement error Args: expert (List[np.ndarray]): all position trajectories for all expert rollouts @@ -24,8 +24,23 @@ def rwse(expert:List[np.ndarray], policy:List[np.ndarray], dt:float=0.1) -> floa assert len(expert) == len(policy) # calculate rwse + rwse = [] + for expert_trajectory, policy_trajectory in zip(expert, policy): + _, T1 = expert_trajectory.shape + _, T2 = policy_trajectory.shape + minT = min(T1, T2) - return 0.0 + crop_expert_trajectory = expert_trajectory[:, :minT] + crop_policy_trajectory = policy_trajectory[:, :minT] + e = ((crop_policy_trajectory - crop_expert_trajectory)**2).sum(0).mean() + + rwse.append(e) + + assert len(rwse) == len(expert) + rwse = np.array(rwse) + avg_rwse = rwse.mean() + + return avg_rwse