-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmetrics.py
More file actions
56 lines (40 loc) · 1.95 KB
/
Copy pathmetrics.py
File metadata and controls
56 lines (40 loc) · 1.95 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
"""Displacement metrics for vessel-trajectory forecasting (metres).
All functions take positions already inverse-normalised to metres in the
local anchor frame. Single-future arrays are shaped ``(N, T, 2)``;
multi-future arrays are shaped ``(N, K, T, 2)``.
"""
from __future__ import annotations
import numpy as np
__all__ = ["ade", "fde", "min_ade", "min_fde", "displacement_table"]
def _check(pred, gt):
pred = np.asarray(pred, float)
gt = np.asarray(gt, float)
if pred.shape != gt.shape:
raise ValueError(f"pred {pred.shape} and gt {gt.shape} must match")
if pred.shape[-1] != 2:
raise ValueError("last dimension must be 2 (x, y)")
return pred, gt
def ade(pred, gt) -> float:
"""Average displacement error over all steps and samples (metres)."""
pred, gt = _check(pred, gt)
return float(np.linalg.norm(pred - gt, axis=-1).mean())
def fde(pred, gt) -> float:
"""Final displacement error at the last predicted step (metres)."""
pred, gt = _check(pred, gt)
return float(np.linalg.norm(pred[..., -1, :] - gt[..., -1, :], axis=-1).mean())
def _best_k(pred, gt, reducer):
"""pred (N,K,T,2), gt (N,T,2) -> per-sample best-of-K error, then mean."""
pred = np.asarray(pred, float)
gt = np.asarray(gt, float)[:, None] # (N,1,T,2)
d = np.linalg.norm(pred - gt, axis=-1) # (N,K,T)
per_k = reducer(d) # (N,K)
return float(per_k.min(axis=1).mean())
def min_ade(pred, gt) -> float:
"""minADE@K: best-of-K average displacement error. pred is (N,K,T,2)."""
return _best_k(pred, gt, lambda d: d.mean(axis=-1))
def min_fde(pred, gt) -> float:
"""minFDE@K: best-of-K final displacement error. pred is (N,K,T,2)."""
return _best_k(pred, gt, lambda d: d[..., -1])
def displacement_table(pred, gt) -> dict:
"""Convenience: return {'ade': ..., 'fde': ...} for a single-future model."""
return {"ade": ade(pred, gt), "fde": fde(pred, gt)}