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