""" MotionTransformerGraphV11 — Post-decoder graph insertion. Current V6: graph after A-attn, BEFORE context fusion + motion decoder. V11: graph AFTER motion decoder, applied to readout_token. Rationale: by the time readout_token is produced, the motion decoder has deeply processed the context. The graph can then refine the final representation with future interaction, closer to the output heads. This means the graph correction directly influences the trajectory output without being diluted by subsequent decoder layers. """ 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 MotionTransformerGraphV11(MotionTransformerGraph): """Graph after motion decoder instead of before context fusion.""" 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) # Post-decoder graph (operates on readout_token dim, not y_emb dim) dim_decoder = model_config.MOTION_DECODER.D_MODEL self.post_graph = FutureInteractionGraphV6( embed_dim=dim_decoder, future_steps=self.T_future, num_agents=self.A, num_heads=4, dropout=graph_dropout, num_gnn_layers=graph_num_gnn_layers, time_dim=dim_decoder, top_n_neighbors=top_n_neighbors, rel_traj_hidden=rel_traj_hidden, y0_score_dim=y0_score_dim) # Time projection to decoder dim (if different from encoder dim) if dim_decoder != self.dim: self.t_emb_proj = nn.Linear(self.dim, dim_decoder) else: self.t_emb_proj = nn.Identity() # Remove parent's pre-decoder graph (set to None to skip in parent forward) # Actually we'll override _forward_impl entirely params = sum(p.numel() for p in self.post_graph.parameters()) logger.info(f"V11: Post-decoder FutureInteractionGraphV6 params: {params:,}") 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) # K-attn 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) # A-attn (no graph here — standard) 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) # Embedding dropout 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.) # NO graph before context fusion # Context fusion + motion decoder (same as baseline) 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) # [B, K, A, D_dec] # ---- POST-DECODER GRAPH ---- if not skip_graph: y_graph_src = (y_0_for_graph.view(B, K, A, self.T_future, 2) if y_0_for_graph is not None else 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_ t_emb_dec = self.t_emb_proj(t_emb) readout_graph = self.post_graph( readout_token, y_abs, t_emb_dec, tau, sigma_agent=sigma_for_graph) readout_token = readout_graph # Output heads 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