File size: 1,844 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
"""
MotionTransformerGraphV4 — two-pass backbone with learnable sparse graph.

Uses FutureInteractionGraphV4: cheap geometric scorer selects top-N neighbors
per agent per mode, then RelTrajEncoder runs only on selected edges.
"""

from models.backbone_graph import MotionTransformerGraph
from models.graph_interaction_nba_v4 import FutureInteractionGraphV4


class MotionTransformerGraphV4(MotionTransformerGraph):
    """MotionTransformerGraph with learnable sparse future interaction graph.

    Extra constructor kwargs:
        top_n_neighbors (int, default 5): neighbors kept per agent per mode.
        rel_traj_hidden (int, default 32): hidden dim in RelTrajEncoder.
    """

    def __init__(self, model_config, logger, config,
                 graph_num_gnn_layers: int = 2,
                 graph_dropout: float = 0.1,
                 top_n_neighbors: int = 5,
                 rel_traj_hidden: int = 32):
        super().__init__(model_config, logger, config,
                         graph_num_gnn_layers=graph_num_gnn_layers,
                         graph_dropout=graph_dropout)

        D        = self.dim
        time_dim = D

        self.future_graph = FutureInteractionGraphV4(
            embed_dim       = D,
            future_steps    = self.T_future,
            num_agents      = self.A,
            num_heads       = 4,
            dropout         = graph_dropout,
            num_gnn_layers  = graph_num_gnn_layers,
            time_dim        = time_dim,
            top_n_neighbors = top_n_neighbors,
            rel_traj_hidden = rel_traj_hidden,
        )

        params_graph = sum(p.numel() for p in self.future_graph.parameters())
        logger.info("FutureInteractionGraphV4 parameters: {:,}".format(params_graph))
        logger.info("Top-N neighbors: {:d} / {:d}".format(top_n_neighbors, self.A - 1))