| """ |
| MotionTransformerGraphV10 — Future-aware attention bias in agent self-attention. |
| |
| Instead of adding a separate graph module, MODULATE the existing A-level |
| self-attention by adding a future-proximity bias to attention logits. |
| |
| When agents i and j have overlapping/close future trajectories, the A-attn |
| should pay more attention to their interaction. |
| |
| Implementation: compute pairwise future min-distance → attention bias matrix |
| → add to noisy_y_attn_a's attention scores. No new GNN, no new modules. |
| Only a lightweight bias computation from y_0_for_graph. |
| """ |
|
|
| import math |
| import torch |
| import torch.nn as nn |
| from einops import rearrange, repeat |
|
|
| from models.backbone_graph import MotionTransformerGraph |
| from models.graph_interaction_nba_v6 import FutureInteractionGraphV6 |
|
|
|
|
| class FutureAttnBias(nn.Module): |
| """Compute pairwise future interaction bias for attention.""" |
|
|
| def __init__(self, hidden_dim, num_heads=4): |
| super().__init__() |
| self.num_heads = num_heads |
| |
| |
| self.bias_proj = nn.Sequential( |
| nn.Linear(3, hidden_dim), |
| nn.ReLU(), |
| nn.Linear(hidden_dim, num_heads), |
| ) |
| |
| self.scale = nn.Parameter(torch.tensor(0.1)) |
|
|
| def forward(self, y_abs): |
| """ |
| Args: |
| y_abs: [B, K, A, T, 2] predicted absolute future positions |
| Returns: |
| bias: [B*K, num_heads, A, A] attention bias |
| """ |
| B, K, A, T, _ = y_abs.shape |
|
|
| |
| pos_i = y_abs.unsqueeze(3) |
| pos_j = y_abs.unsqueeze(2) |
| dist = (pos_i - pos_j).norm(dim=-1) |
|
|
| |
| min_dist = dist.min(dim=-1).values |
| mean_dist = dist.mean(dim=-1) |
|
|
| |
| convergence = dist[..., 0] - dist[..., -1] |
|
|
| feat = torch.stack([min_dist, mean_dist, convergence], dim=-1) |
|
|
| |
| bias = self.bias_proj(feat) |
| bias = bias.permute(0, 1, 4, 2, 3) |
| bias = bias.reshape(B * K, self.num_heads, A, A) |
|
|
| return self.scale * bias |
|
|
|
|
| class FutureAwareAttnA(nn.Module): |
| """Agent self-attention with future interaction bias. |
| |
| Wraps nn.TransformerEncoderLayer, adding a future-based bias to attention. |
| """ |
|
|
| def __init__(self, d_model, nhead, dim_feedforward, dropout): |
| super().__init__() |
| self.base_attn = nn.TransformerEncoderLayer( |
| d_model=d_model, nhead=nhead, |
| dim_feedforward=dim_feedforward, dropout=dropout, |
| batch_first=True) |
| self.nhead = nhead |
| self.d_k = d_model // nhead |
| self._bias = None |
|
|
| def set_future_bias(self, bias): |
| """Set the future attention bias [B*K, nhead, A, A].""" |
| self._bias = bias |
|
|
| def forward(self, src, src_mask=None, src_key_padding_mask=None, is_causal=False): |
| if self._bias is None: |
| return self.base_attn(src, src_mask, src_key_padding_mask) |
|
|
| |
| |
| x = self.base_attn.norm1(src) |
|
|
| |
| qkv_weight = self.base_attn.self_attn.in_proj_weight |
| qkv_bias_param = self.base_attn.self_attn.in_proj_bias |
|
|
| B_K, A, D = x.shape |
| qkv = torch.nn.functional.linear(x, qkv_weight, qkv_bias_param) |
| q, k, v = qkv.chunk(3, dim=-1) |
|
|
| q = q.view(B_K, A, self.nhead, self.d_k).permute(0, 2, 1, 3) |
| k = k.view(B_K, A, self.nhead, self.d_k).permute(0, 2, 1, 3) |
| v = v.view(B_K, A, self.nhead, self.d_k).permute(0, 2, 1, 3) |
|
|
| |
| attn = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k) |
| attn = attn + self._bias |
| attn = torch.softmax(attn, dim=-1) |
| attn = torch.nn.functional.dropout(attn, p=self.base_attn.self_attn.dropout, training=self.training) |
|
|
| out = torch.matmul(attn, v) |
| out = out.permute(0, 2, 1, 3).reshape(B_K, A, D) |
| out = self.base_attn.self_attn.out_proj(out) |
|
|
| |
| src2 = self.base_attn.dropout1(out) |
| src = src + src2 |
| src = self.base_attn.norm2(src) |
| src2 = self.base_attn.linear2( |
| self.base_attn.dropout(self.base_attn.activation(self.base_attn.linear1(src)))) |
| src = src + self.base_attn.dropout2(src2) |
|
|
| self._bias = None |
| return src |
|
|
|
|
| class MotionTransformerGraphV10(MotionTransformerGraph): |
| """Backbone with future-aware attention bias in A-attn + V6 graph.""" |
|
|
| def __init__(self, model_config, logger, config, |
| graph_num_gnn_layers=2, graph_dropout=0.1, |
| top_n_neighbors=5, rel_traj_hidden=32, y0_score_dim=32): |
| super().__init__(model_config, logger, config, |
| graph_num_gnn_layers=graph_num_gnn_layers, |
| graph_dropout=graph_dropout) |
|
|
| |
| self.future_graph = FutureInteractionGraphV6( |
| embed_dim=self.dim, future_steps=self.T_future, |
| num_agents=self.A, num_heads=4, dropout=graph_dropout, |
| num_gnn_layers=graph_num_gnn_layers, time_dim=self.dim, |
| top_n_neighbors=top_n_neighbors, |
| rel_traj_hidden=rel_traj_hidden, y0_score_dim=y0_score_dim) |
|
|
| |
| self.future_bias = FutureAttnBias(hidden_dim=32, num_heads=4) |
|
|
| |
| dropout_ = model_config.MOTION_DECODER.DROPOUT_OF_ATTN |
| self.noisy_y_attn_a = FutureAwareAttnA( |
| d_model=self.dim, nhead=4, |
| dim_feedforward=self.dim * 4, dropout=dropout_) |
|
|
| params_bias = sum(p.numel() for p in self.future_bias.parameters()) |
| params_graph = sum(p.numel() for p in self.future_graph.parameters()) |
| logger.info(f"V10: FutureAttnBias params: {params_bias:,}") |
| logger.info(f"V10: FutureInteractionGraphV6 params: {params_graph:,}") |
|
|
| def _forward_impl(self, y, time, x_data, |
| y_0_for_graph=None, sigma_for_graph=None, |
| skip_graph=False): |
| |
| if y.size(-1) == 2: |
| y = y.reshape((-1, self.model_cfg.NUM_PROPOSED_QUERY, |
| self.A, self.T_future * 2)) |
| device = y.device |
| B, K, A, _ = y.shape |
|
|
| |
| encoder_out = self.context_encoder(x_data['past_traj_original_scale']) |
| encoder_out_batch = repeat(encoder_out, 'b a d -> b k a d', k=K, a=A) |
|
|
| |
| y_emb = self.noisy_y_mlp(y) |
|
|
| |
| time_ = time |
| if self.config.denoising_method == 'fm': |
| time = time * 1000.0 |
| t_emb = self.time_mlp(time) |
| t_emb_batch = repeat(t_emb, 'b d -> b k a d', b=B, k=K, a=A) |
|
|
| |
| k_pe = self.motion_query_embedding( |
| torch.arange(self.model_cfg.NUM_PROPOSED_QUERY, device=device)) |
| k_pe_batch = repeat(k_pe, 'k d -> b k a d', b=B, a=A) |
| a_pe = self.agent_order_embedding( |
| torch.arange(self.model_cfg.CONTEXT_ENCODER.NUM_OF_ATTN_NEIGHBORS, device=device)) |
| a_pe_batch = repeat(a_pe, 'a d -> b k a d', b=B, k=K) |
|
|
| |
| y_emb_k = rearrange(self.apply_PE(y_emb, k_pe_batch, a_pe_batch), 'b k a d -> (b a) k d') |
| y_emb_k = self.noisy_y_attn_k(y_emb_k) |
| y_emb = rearrange(y_emb_k, '(b a) k d -> b k a d', b=B, a=A) |
|
|
| |
| if not skip_graph and y_0_for_graph is not None: |
| y_graph_src = y_0_for_graph.view(B, K, A, self.T_future, 2) |
| y_graph_unnorm = self._unnormalize_y(y_graph_src) |
| init_pos = x_data['past_traj_original_scale'][:, :, -1, :2] |
| y_abs = y_graph_unnorm + init_pos.unsqueeze(1).unsqueeze(3) |
|
|
| |
| fut_bias = self.future_bias(y_abs) |
| self.noisy_y_attn_a.set_future_bias(fut_bias) |
|
|
| y_emb_a = rearrange(y_emb, 'b k a d -> (b k) a d') |
| y_emb_a = self.noisy_y_attn_a(y_emb_a) |
| y_emb = rearrange(y_emb_a, '(b k) a d -> b k a d', b=B, k=K) |
|
|
| |
| if self.training and self.config.get('drop_method', None) == 'emb': |
| m, k_drop = self.config.drop_logi_m, self.config.drop_logi_k |
| p_m = 1 / (1 + torch.exp(-k_drop * (time_ - m))) |
| p_m = p_m[:, None, None, None] |
| y_emb = y_emb.masked_fill(torch.rand_like(p_m) < p_m, 0.) |
|
|
| |
| if not skip_graph: |
| if y_abs is None: |
| y_graph_src = y.view(B, K, A, self.T_future, 2) |
| y_graph_unnorm = self._unnormalize_y(y_graph_src) |
| init_pos = x_data['past_traj_original_scale'][:, :, -1, :2] |
| y_abs = y_graph_unnorm + init_pos.unsqueeze(1).unsqueeze(3) |
|
|
| tau = time_ |
| y_emb_graph = self.future_graph( |
| y_emb, y_abs, t_emb, tau, sigma_agent=sigma_for_graph) |
| y_emb = y_emb_graph |
|
|
| |
| emb_fusion = self.init_emb_fusion_mlp( |
| torch.cat((encoder_out_batch, y_emb, t_emb_batch), dim=-1)) |
| query_token = self.post_pe_cat_mlp( |
| self.apply_PE(emb_fusion, k_pe_batch, a_pe_batch)) |
| readout_token = self.motion_decoder(query_token, t_emb) |
|
|
| denoiser_x = self.reg_head(readout_token) |
| denoiser_cls = self.cls_head(readout_token).squeeze(-1) |
| logvar = self.logvar_head(readout_token) |
|
|
| return denoiser_x, denoiser_cls, logvar |
|
|