| """ |
| FutureInteractionGraphV5 β uncertainty-aware semantic scorer + sparse RelTrajEncoder. |
| |
| Builds on V4 (learnable top-N sparse graph) with an enhanced edge scorer: |
| |
| V4 scorer: [mean_rel, std_rel, min_dist] β [E, 5] (geometric only) |
| V5 scorer: [y0_emb_i, y0_emb_j_scaled, β [E, 2*D_score + 7] |
| Ο_i, Ο_j, |
| mean_rel, std_rel, min_dist] |
| |
| Where: |
| y0_emb_{i,j} β trajectory embedding from clean y_0_hat (not noisy y_t). |
| Projected from y_abs [B,K,A,T,2] which is already y_0_hat |
| unnormalised β no noise, same source as edge geometry. |
| Ο_i, Ο_j β explicit per-agent mean uncertainty, allowing the scorer to |
| learn "uncertain target prefers certain source." |
| |
| When sigma_agent is None (pass 1 of the two-pass forward), certainty scaling |
| is skipped and Ο features are zeroed β scorer still runs on geometric + |
| semantic features, gracefully degrading to a noisier but functional signal. |
| """ |
|
|
| import torch |
| import torch.nn as nn |
| from models.graph_interaction_nba_v4 import FutureInteractionGraphV4 |
| from models.graph_interaction_nba_v3 import RelTrajEncoder |
|
|
|
|
| def _heading_diff(pos_i: torch.Tensor, |
| pos_j: torch.Tensor) -> torch.Tensor: |
| """Per-timestep relative heading angle difference. |
| |
| Args: |
| pos_i, pos_j: [E, T, 2] absolute positions of target and source |
| Returns: |
| [E, T, 1] angle difference wrapped to [-Ο, Ο] |
| β positive: j is turning more CCW than i |
| β β 0: parallel motion; β Β±Ο: opposing motion |
| """ |
| vel_i = pos_i[:, 1:] - pos_i[:, :-1] |
| vel_j = pos_j[:, 1:] - pos_j[:, :-1] |
|
|
| h_i = torch.atan2(vel_i[..., 1], vel_i[..., 0]) |
| h_j = torch.atan2(vel_j[..., 1], vel_j[..., 0]) |
|
|
| diff = torch.atan2(torch.sin(h_j - h_i), |
| torch.cos(h_j - h_i)) |
|
|
| |
| diff = torch.cat([diff, diff[:, -1:]], dim=1) |
| return diff.unsqueeze(-1) |
|
|
|
|
| class FutureInteractionGraphV5(FutureInteractionGraphV4): |
| """Sparse interaction graph with uncertainty-aware semantic edge scorer. |
| |
| Extra constructor kwarg (beyond V4): |
| y0_score_dim (int, default 32): projection dim for y_0_hat trajectory |
| embedding used in the scorer. |
| """ |
|
|
| def __init__(self, embed_dim: int, future_steps: int, num_agents: int, |
| num_heads: int = 4, dropout: float = 0.1, |
| num_gnn_layers: int = 2, time_dim: int = 128, |
| top_n_neighbors: int = 5, rel_traj_hidden: int = 32, |
| y0_score_dim: int = 32): |
| super().__init__( |
| embed_dim = embed_dim, |
| future_steps = future_steps, |
| num_agents = num_agents, |
| num_heads = num_heads, |
| dropout = dropout, |
| num_gnn_layers = num_gnn_layers, |
| time_dim = time_dim, |
| top_n_neighbors = top_n_neighbors, |
| rel_traj_hidden = rel_traj_hidden, |
| ) |
| self.y0_score_dim = y0_score_dim |
|
|
| |
| |
| self.rel_traj_encoder = RelTrajEncoder( |
| out_dim = embed_dim, |
| T = future_steps, |
| D_hidden = rel_traj_hidden, |
| num_heads = 4, |
| in_channels = 3, |
| ) |
|
|
| |
| |
| self.y0_score_proj = nn.Sequential( |
| nn.Linear(future_steps * 2, y0_score_dim), |
| nn.ReLU(inplace=True), |
| ) |
|
|
| |
| |
| |
| self.score_proj_i = nn.Linear(y0_score_dim, y0_score_dim) |
| self.score_proj_j = nn.Linear(y0_score_dim, y0_score_dim) |
| |
| |
| |
| |
| |
| |
| |
| |
| head_in = y0_score_dim * 2 + 7 |
| self.edge_scorer = nn.Sequential( |
| nn.Linear(head_in, 32), |
| nn.ReLU(inplace=True), |
| nn.Linear(32, 1), |
| ) |
|
|
| |
| |
| |
|
|
| def forward( |
| self, |
| y_emb: torch.Tensor, |
| y_abs: torch.Tensor, |
| t_emb: torch.Tensor, |
| tau: torch.Tensor, |
| sigma_agent: torch.Tensor = None, |
| ) -> torch.Tensor: |
| B, K, A, D = y_emb.shape |
| T = y_abs.shape[3] |
| E0 = self._E0 |
| N = self.top_n |
|
|
| |
| pos_bk = y_abs.reshape(B * K * A, T, 2) |
| edge_index_bk = self._make_batched_edge_index(B * K) |
|
|
| pos_i_t = pos_bk[edge_index_bk[1]] |
| pos_j_t = pos_bk[edge_index_bk[0]] |
| rel_pos_t = pos_j_t - pos_i_t |
|
|
| |
| mean_rel = rel_pos_t.mean(dim=1) |
| std_rel = rel_pos_t.std(dim=1) |
| min_dist = (rel_pos_t.norm(dim=-1) |
| .min(dim=1).values |
| .unsqueeze(-1)) |
|
|
| |
| y0_flat = y_abs.reshape(B * K * A, T * 2) |
| y0_emb = self.y0_score_proj(y0_flat) |
|
|
| y0_emb_i = y0_emb[edge_index_bk[1]] |
| y0_emb_j = y0_emb[edge_index_bk[0]] |
|
|
| |
| if sigma_agent is not None: |
| sigma_mean = sigma_agent.mean(dim=-1) |
| sigma_bka = sigma_mean.reshape(B * K * A) |
| sigma_i = sigma_bka[edge_index_bk[1]].unsqueeze(-1) |
| sigma_j = sigma_bka[edge_index_bk[0]].unsqueeze(-1) |
| tau_bka = sigma_bka |
| else: |
| |
| sigma_i = torch.zeros(rel_pos_t.size(0), 1, device=y_abs.device) |
| sigma_j = torch.zeros_like(sigma_i) |
| tau_bka = (tau |
| .unsqueeze(1).unsqueeze(2) |
| .expand(-1, K, A) |
| .reshape(B * K * A)) |
|
|
| |
| |
| h_i = self.score_proj_i(y0_emb_i) |
| h_j = self.score_proj_j(y0_emb_j) |
| interact = torch.cat([h_i * h_j, |
| h_i - h_j], dim=-1) |
|
|
| score_feat = torch.cat([ |
| interact, |
| sigma_i, |
| sigma_j, |
| mean_rel, |
| std_rel, |
| min_dist, |
| ], dim=-1) |
| scores = self.edge_scorer(score_feat).squeeze(-1) |
|
|
| |
| scores_grouped = scores.view(B * K * A, A - 1) |
| _, top_idx = scores_grouped.topk(N, dim=-1, sorted=False) |
| mask = torch.zeros(B * K * A, A - 1, |
| device=scores.device, dtype=torch.bool) |
| mask.scatter_(1, top_idx, True) |
| mask_flat = mask.view(-1) |
|
|
| |
| |
| heading = _heading_diff(pos_i_t[mask_flat], |
| pos_j_t[mask_flat]) |
| rel_pos_sparse = torch.cat( |
| [rel_pos_t[mask_flat], heading], dim=-1 |
| ) |
|
|
| if sigma_agent is not None: |
| sigma_full = sigma_agent.reshape(B * K * A, T) |
| sigma_i_t = sigma_full[edge_index_bk[1][mask_flat]] |
| sigma_j_t = sigma_full[edge_index_bk[0][mask_flat]] |
| sigma_bias = sigma_i_t - sigma_j_t |
| else: |
| sigma_bias = None |
|
|
| edge_attr_sparse = self.rel_traj_encoder( |
| rel_pos_sparse, sigma_bias |
| ) |
|
|
| |
| edge_index_sparse = edge_index_bk[:, mask_flat] |
|
|
| temb_bka = (t_emb |
| .unsqueeze(1).unsqueeze(2) |
| .expand(-1, K, A, -1) |
| .reshape(B * K * A, D)) |
|
|
| nodes = y_emb.reshape(B * K * A, D) |
| for layer in self.gnn_layers: |
| nodes = layer(nodes, edge_index_sparse, edge_attr_sparse, |
| temb_agent=temb_bka, tau=tau_bka) |
|
|
| |
| orig = y_emb.reshape(B * K * A, D) |
| gate = self.gate_proj(torch.cat([orig, nodes], dim=-1)) |
| out = orig + gate * self.out_proj(nodes) |
|
|
| return out.view(B, K, A, D) |
|
|