File size: 6,744 Bytes
d4cbafd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 | """
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
|