sra-trajectory-code / MoFlow /models /backbone_graph_v4.py
po03087's picture
SRA: MID/LED/MoFlow code + RUNNING.md instructions (code only, no data/ckpts)
d4cbafd verified
Raw
History Blame Contribute Delete
1.84 kB
"""
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))