File size: 5,590 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
"""
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