sra-trajectory-code / MoFlow /models /backbone_graph_v14.py
po03087's picture
SRA: MID/LED/MoFlow code + RUNNING.md instructions (code only, no data/ckpts)
d4cbafd verified
Raw
History Blame Contribute Delete
6.74 kB
"""
MotionTransformerGraphV14 — V6 graph + collision-aware auxiliary loss.
Same V6 architecture, but adds a penalty for predicted trajectories where
agents come too close (potential collisions). This doesn't change the graph
module — it adds a TRAINING SIGNAL that encourages collision avoidance.
The collision loss is computed from the reg_head output (predicted trajectories)
and added as a 5th return value. The trainer adds it to the total loss.
"""
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
def _collision_loss(denoiser_x, init_pos, T, A, unnorm_fn=None,
threshold=1.0, margin=0.5):
"""Penalize predicted trajectories where agents come too close.
Args:
denoiser_x: [B, K, A, T*2] predicted trajectory (may be normalized)
init_pos: [B, A, 2] last observed positions
T: future timesteps
A: num agents
threshold: distance below which we penalize (meters)
margin: soft margin for penalty
Returns:
scalar loss
"""
B, K = denoiser_x.shape[:2]
# Reshape to [B, K, A, T, 2]
traj = denoiser_x.view(B, K, A, T, 2)
# Convert to absolute positions
# traj is relative to last obs in unnormalized space
if unnorm_fn is not None:
traj = unnorm_fn(traj)
traj_abs = traj + init_pos[:, None, :, None, :] # [B, K, A, T, 2]
# Pairwise distances: [B, K, A, A, T]
pos_i = traj_abs.unsqueeze(3) # [B, K, A, 1, T, 2]
pos_j = traj_abs.unsqueeze(2) # [B, K, 1, A, T, 2]
dist = (pos_i - pos_j).norm(dim=-1) # [B, K, A, A, T]
# Mask self-loops
eye = torch.eye(A, device=dist.device).bool()
dist = dist.masked_fill(eye[None, None, :, :, None], float('inf'))
# Min distance per pair over time
min_dist = dist.min(dim=-1).values # [B, K, A, A]
# Soft penalty: max(0, threshold - min_dist + margin) for pairs that are too close
penalty = torch.relu(threshold - min_dist + margin)
# Average over all pairs, modes, and batch
loss = penalty.sum() / (B * K * A * (A - 1))
return loss
class MotionTransformerGraphV14(MotionTransformerGraph):
"""V6 graph + collision-aware auxiliary loss."""
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.collision_weight = config.get('collision_weight', 0.1)
p = sum(p.numel() for p in self.future_graph.parameters())
logger.info(f"V14: Graph params: {p:,}, collision_weight: {self.collision_weight}")
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)
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:
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_
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)
# Store collision loss info for the flow_matching wrapper
if self.training:
init_pos = x_data['past_traj_original_scale'][:, :, -1, :2]
self._collision_loss = self.collision_weight * _collision_loss(
denoiser_x.detach() if not self.training else denoiser_x,
init_pos, self.T_future, A,
unnorm_fn=self._unnormalize_y)
else:
self._collision_loss = torch.tensor(0.0, device=device)
return denoiser_x, denoiser_cls, logvar