| """ |
| MotionTransformerGraphV6 — two-pass backbone with RAG-style sparse graph. |
| """ |
|
|
| from models.backbone_graph import MotionTransformerGraph |
| from models.graph_interaction_nba_v6 import FutureInteractionGraphV6 |
| from models.interaction_baselines import build_interaction_module |
|
|
|
|
| class MotionTransformerGraphV6(MotionTransformerGraph): |
| """MotionTransformerGraph with RAG-style sparse future interaction graph. |
| |
| Extra constructor kwargs: |
| top_n_neighbors (int, default 5) |
| rel_traj_hidden (int, default 32) |
| y0_score_dim (int, default 32) |
| """ |
|
|
| 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, |
| y0_score_dim: int = 32, |
| edge_mode: str = 'full', |
| neighbor_mode: str = 'rag'): |
| super().__init__(model_config, logger, config, |
| graph_num_gnn_layers=graph_num_gnn_layers, |
| graph_dropout=graph_dropout) |
|
|
| self.future_graph = build_interaction_module( |
| 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, |
| edge_mode = edge_mode, |
| neighbor_mode = neighbor_mode, |
| ) |
|
|
| params_graph = sum(p.numel() for p in self.future_graph.parameters()) |
| logger.info("FutureInteractionGraphV6 parameters: {:,}".format(params_graph)) |
| logger.info("Top-N neighbors: {:d} / {:d} edge_mode: {} neighbor_mode: {}".format( |
| top_n_neighbors, self.A - 1, edge_mode, neighbor_mode)) |
|
|