sra-trajectory-code / MoFlow /models /backbone_graph_v11.py
po03087's picture
SRA: MID/LED/MoFlow code + RUNNING.md instructions (code only, no data/ckpts)
d4cbafd verified
Raw
History Blame Contribute Delete
5.59 kB
"""
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