""" 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 # E4: SRA_MODULE factory 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( # E4: sra|gameformer|c2f via env SRA_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))