| """ |
| FutureInteractionGraphV6 β RAG-style neighbor retrieval with sparse RelTrajEncoder. |
| |
| Replaces V5's per-edge MLP scorer with a RAG-style query-key dot product: |
| |
| Per-node (55K, computed once): |
| q_i = W_q([y0_emb_i, Ο_i]) β what agent i is looking for |
| k_j = W_k([y0_emb_j, Ο_j]) β what agent j offers |
| |
| Per-edge (550K, cheap): |
| semantic_score = (q_i Β· k_j) / βD_s |
| geo_bias = geo_mlp([mean_rel, std_rel, min_dist, heading_diff_mean]) |
| score = semantic_score + geo_bias |
| |
| Top-N selection β RelTrajEncoder([rel_pos, heading_diff]) on selected edges. |
| |
| Design rationale: |
| - Asymmetric W_q / W_k: "what to look for" β "how to present oneself" |
| - Uncertainty in query/key: uncertain agents learn to seek certain neighbors |
| through training, not hard-coded scaling |
| - Geometric features as additive re-ranking bias: cleanly separates |
| semantic retrieval from spatial prior |
| - Heading in geo_bias: converging agents are more likely to interact |
| - Per-node query/key is ~10x cheaper than per-edge MLP scorer in V5 |
| |
| Document content (unchanged from V5): |
| RelTrajEncoder([rel_pos(2), heading_diff(1)]) β [E_selected, D] |
| β GNN message passing on sparse graph. |
| """ |
|
|
| import os |
| import torch |
| import torch.nn as nn |
| from models.graph_interaction_nba_v4 import FutureInteractionGraphV4 |
| from models.graph_interaction_nba_v3 import RelTrajEncoder |
| from models.graph_interaction_nba_v5 import _heading_diff |
|
|
|
|
| class FutureInteractionGraphV6(FutureInteractionGraphV4): |
| """RAG-style sparse interaction graph. |
| |
| Extra constructor kwargs (beyond V4): |
| y0_score_dim (int, default 32): dim of query/key embedding space. |
| """ |
|
|
| EDGE_MODES = { |
| 'full': 3, |
| 'dist_only': 1, |
| 'relpos_only': 2, |
| 'heading_only': 1, |
| 'full_relvel': 5, |
| 'vel_only': 2, |
| } |
|
|
| 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, edge_mode: str = 'full', |
| neighbor_mode: str = 'rag'): |
| 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, |
| ) |
| assert edge_mode in self.EDGE_MODES, f"Unknown edge_mode: {edge_mode}" |
| self.edge_mode = edge_mode |
| in_ch = self.EDGE_MODES[edge_mode] |
| assert neighbor_mode in ('rag', 'l2', 'semantic'), f"Unknown neighbor_mode: {neighbor_mode}" |
| self.neighbor_mode = neighbor_mode |
|
|
| self.y0_score_dim = y0_score_dim |
| self.scale = y0_score_dim ** -0.5 |
| self._original_top_n = top_n_neighbors |
|
|
| |
| self.y0_score_proj = nn.Sequential( |
| nn.Linear(future_steps * 2, y0_score_dim), |
| nn.ReLU(inplace=True), |
| ) |
|
|
| |
| self.W_q = nn.Linear(y0_score_dim + 1, y0_score_dim) |
| self.W_k = nn.Linear(y0_score_dim + 1, y0_score_dim) |
|
|
| |
| |
| self.geo_mlp = nn.Sequential( |
| nn.Linear(6, 16), |
| nn.ReLU(inplace=True), |
| nn.Linear(16, 1), |
| ) |
|
|
| del self.edge_scorer |
|
|
| self.rel_traj_encoder = RelTrajEncoder( |
| out_dim = embed_dim, |
| T = future_steps, |
| D_hidden = rel_traj_hidden, |
| num_heads = 4, |
| in_channels = in_ch, |
| ) |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| if os.environ.get('SRA_SOFT_START', '') not in ('', '0', 'false', 'False'): |
| nn.init.zeros_(self.out_proj.weight) |
| nn.init.zeros_(self.out_proj.bias) |
| _gb = float(os.environ.get('SRA_GATE_BIAS', -4.0) or -4.0) |
| for _m in self.gate_proj.modules(): |
| if isinstance(_m, nn.Linear): |
| nn.init.constant_(_m.bias, _gb) |
|
|
| |
| |
| |
|
|
| def forward( |
| self, |
| y_emb: torch.Tensor, |
| y_abs: torch.Tensor, |
| t_emb: torch.Tensor, |
| tau: torch.Tensor, |
| sigma_agent: torch.Tensor = None, |
| agent_mask: torch.Tensor = None, |
| ) -> torch.Tensor: |
| B, K, A, D = y_emb.shape |
| T = y_abs.shape[3] |
|
|
| |
| |
| if A <= 1: |
| return y_emb |
|
|
| |
| |
| |
| if A != self.num_agents: |
| self.num_agents = A |
| self._E0 = A * (A - 1) |
| self.top_n = max(1, min(self._original_top_n, A - 1)) |
| src, dst = [], [] |
| for i in range(A): |
| for j in range(A): |
| if i != j: |
| src.append(j); dst.append(i) |
| self._single_edge_index = torch.tensor( |
| [src, dst], dtype=torch.long, device=y_emb.device |
| ) |
|
|
| E0 = self._E0 |
| N = self.top_n |
|
|
| |
| y0_flat = y_abs.reshape(B * K * A, T * 2) |
| y0_emb = self.y0_score_proj(y0_flat) |
|
|
| |
| if sigma_agent is not None: |
| sigma_mean = sigma_agent.mean(dim=-1).reshape(B * K * A, 1) |
| tau_bka = sigma_mean.squeeze(-1) |
| else: |
| sigma_mean = torch.zeros(B * K * A, 1, device=y_abs.device) |
| tau_bka = (tau |
| .unsqueeze(1).unsqueeze(2) |
| .expand(-1, K, A) |
| .reshape(B * K * A)) |
|
|
| |
| node_feat = torch.cat([y0_emb, sigma_mean], dim=-1) |
| q_bka = self.W_q(node_feat) |
| k_bka = self.W_k(node_feat) |
|
|
| |
| 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) |
|
|
| |
| heading_full = _heading_diff(pos_i_t, pos_j_t) |
| heading_mean = heading_full.mean(dim=1) |
|
|
| |
| if self.neighbor_mode == 'l2': |
| |
| avg_dist = rel_pos_t.norm(dim=-1).mean(dim=-1) |
| scores = -avg_dist |
| elif self.neighbor_mode == 'semantic': |
| |
| q_i = q_bka[edge_index_bk[1]] |
| k_j = k_bka[edge_index_bk[0]] |
| scores = (q_i * k_j).sum(dim=-1) * self.scale |
| else: |
| |
| q_i = q_bka[edge_index_bk[1]] |
| k_j = k_bka[edge_index_bk[0]] |
| semantic_score = (q_i * k_j).sum(dim=-1) * self.scale |
|
|
| geo_feat = torch.cat([mean_rel, std_rel, |
| min_dist, heading_mean], dim=-1) |
| geo_bias = self.geo_mlp(geo_feat).squeeze(-1) |
|
|
| scores = semantic_score + geo_bias |
|
|
| |
| if agent_mask is not None: |
| |
| |
| node_real = agent_mask.unsqueeze(1).expand(B, K, A).reshape(B * K * A) |
| |
| src_real = node_real[edge_index_bk[0]] |
| dst_real = node_real[edge_index_bk[1]] |
| edge_real = src_real & dst_real |
| scores = scores.masked_fill(~edge_real, float('-inf')) |
|
|
| |
| scores_grouped = scores.view(B * K * A, A - 1) |
| |
| N_use = min(N, A - 1) |
| _, top_idx = scores_grouped.topk(N_use, 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) |
|
|
| |
| |
| |
| if agent_mask is not None: |
| mask_flat = mask_flat & edge_real |
|
|
| |
| rel_pos_sel = rel_pos_t[mask_flat] |
| heading_sel = heading_full[mask_flat] |
|
|
| if self.edge_mode == 'full': |
| encoder_input = torch.cat([rel_pos_sel, heading_sel], dim=-1) |
| elif self.edge_mode == 'dist_only': |
| encoder_input = rel_pos_sel.norm(dim=-1, keepdim=True) |
| elif self.edge_mode == 'relpos_only': |
| encoder_input = rel_pos_sel |
| elif self.edge_mode == 'heading_only': |
| encoder_input = heading_sel |
| elif self.edge_mode == 'full_relvel': |
| rel_vel = rel_pos_sel[:, 1:] - rel_pos_sel[:, :-1] |
| rel_vel = torch.cat([rel_vel, rel_vel[:, -1:]], dim=1) |
| encoder_input = torch.cat([rel_pos_sel, heading_sel, rel_vel], dim=-1) |
| elif self.edge_mode == 'vel_only': |
| rel_vel = rel_pos_sel[:, 1:] - rel_pos_sel[:, :-1] |
| rel_vel = torch.cat([rel_vel, rel_vel[:, -1:]], dim=1) |
| encoder_input = rel_vel |
|
|
| 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( |
| encoder_input, 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)) |
| |
| _gs = float(os.environ.get('SRA_GATE_SCALE', 1.0) or 1.0) |
| res = _gs * gate * self.out_proj(nodes) |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| _cap = float(os.environ.get('SRA_RES_CAP', 0.0) or 0.0) |
| if _cap > 0: |
| rn = res.norm(dim=-1, keepdim=True) |
| res = res * torch.where(rn > _cap, _cap / rn.clamp_min(1e-6), |
| torch.ones_like(rn)) |
|
|
| |
| |
| |
| |
| |
| _rel = float(os.environ.get('SRA_RES_CAP_REL', 0.0) or 0.0) |
| if _rel > 0: |
| lim = _rel * orig.norm(dim=-1, keepdim=True) |
| rn = res.norm(dim=-1, keepdim=True) |
| res = res * torch.where(rn > lim, lim / rn.clamp_min(1e-6), |
| torch.ones_like(rn)) |
| out = orig + res |
|
|
| return out.view(B, K, A, D) |
|
|