File size: 5,671 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
"""
MotionTransformerGraphV12 — Dual graph insertion: before AND after motion decoder.

Pre-decoder graph enriches y_emb (same as V6).
Post-decoder graph refines readout_token (same as V11).
Both use V6-style interaction graph but separate instances.
"""

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 MotionTransformerGraphV12(MotionTransformerGraph):
    """Dual graph: pre-decoder + post-decoder."""

    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)

        # Pre-decoder graph (enriches y_emb, same as V6)
        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)

        # Post-decoder graph (refines readout_token)
        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=1,  # lighter: 1 GNN layer (refinement only)
            time_dim=dim_decoder,
            top_n_neighbors=top_n_neighbors,
            rel_traj_hidden=rel_traj_hidden, y0_score_dim=y0_score_dim)

        if dim_decoder != self.dim:
            self.t_emb_proj = nn.Linear(self.dim, dim_decoder)
        else:
            self.t_emb_proj = nn.Identity()

        p1 = sum(p.numel() for p in self.future_graph.parameters())
        p2 = sum(p.numel() for p in self.post_graph.parameters())
        logger.info(f"V12: Pre-decoder graph: {p1:,}, Post-decoder graph: {p2:,}")

    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.)

        # Compute geometry once
        y_abs = None
        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_

            # ---- PRE-DECODER GRAPH ----
            y_emb_graph = self.future_graph(
                y_emb, y_abs, t_emb, tau, sigma_agent=sigma_for_graph)
            y_emb = y_emb_graph

        # Context fusion + motion decoder
        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)

        # ---- POST-DECODER GRAPH ----
        if not skip_graph and y_abs is not None:
            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

        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