""" fm_nba_graph_v6.py — Flow-Matching training for NBA with RAG-style sparse Future Interaction Graph (FutureInteractionGraphV6). Difference from fm_nba_graph_v5.py: Replaces the per-edge MLP scorer (V5) with a RAG-style query-key dot product: Per-node (55K, computed once): q_i = W_q([y0_emb_i, σ_i]) — what agent i is looking for k_j = W_k([y0_emb_j, σ_j]) — what agent j offers Per-edge (550K, cheap): semantic_score = (q_i · k_j) / √D_s geo_bias = geo_mlp([mean_rel, std_rel, min_dist, heading_diff_mean]) score = semantic_score + geo_bias Top-N selection → RelTrajEncoder([rel_pos, heading_diff]) on selected edges. Extra CLI flag: --y0_score_dim (default 32): dim of query/key embedding space. Usage: python fm_nba_graph_v6.py --cfg cfg/nba/cor_fm.yml [other flags] """ import os import copy import torch import argparse from torch.utils.data import DataLoader from tensorboardX import SummaryWriter from data.dataloader_nba_graph import NBADatasetMinMax, seq_collate_nba_graph from utils.config import Config from utils.utils import back_up_code_git, set_random_seed, log_config_to_file from models.flow_matching import FlowMatcher from models.backbone_graph_v16 import MotionTransformerGraphV16 from trainer.denoising_model_trainers import Trainer # --------------------------------------------------------------------------- # Argument parsing # --------------------------------------------------------------------------- def parse_config(): parser = argparse.ArgumentParser() # Basic configuration parser.add_argument('--cfg', default='cfg/nba/cor_fm.yml', type=str) parser.add_argument('--exp', default='', type=str) # Data configuration parser.add_argument('--epochs', default=None, type=int) parser.add_argument('--batch_size', default=None, type=int) parser.add_argument('--data_dir', type=str, default='./data/nba') parser.add_argument('--overfit', default=False, action='store_true') parser.add_argument('--n_train', type=int, default=32500) parser.add_argument('--n_test', type=int, default=12500) parser.add_argument('--rotate', default=False, action='store_true') parser.add_argument('--checkpt_freq', default=1, type=int) parser.add_argument('--max_num_ckpts', default=5, type=int) parser.add_argument('--data_norm', default='min_max', choices=['min_max', 'sqrt']) # Reproducibility parser.add_argument('--fix_random_seed', action='store_true', default=False) parser.add_argument('--seed', type=int, default=42) # FM parameters parser.add_argument('--sampling_steps', type=int, default=10) parser.add_argument('--t_schedule', type=str, choices=['uniform', 'logit_normal'], default='logit_normal') parser.add_argument('--fm_skewed_t', default=None, type=str) parser.add_argument('--logit_norm_mean', default=-0.5, type=float) parser.add_argument('--logit_norm_std', default=1.5, type=float) parser.add_argument('--fm_wrapper', type=str, default='direct', choices=['direct', 'velocity', 'precond']) parser.add_argument('--fm_rew_sqrt', default=False, action='store_true') parser.add_argument('--fm_in_scaling', default=False, action='store_true') # Input dropout / masking parser.add_argument('--drop_method', default='emb', type=str, choices=['None', 'input', 'emb']) parser.add_argument('--drop_logi_k', default=20.0, type=float) parser.add_argument('--drop_logi_m', default=0.5, type=float) # Architecture parser.add_argument('--use_pre_norm', default=False, action='store_true') # General denoising parser.add_argument('--tied_noise', default=False, action='store_true') # Loss parser.add_argument('--loss_nn_mode', type=str, default='agent', choices=['agent', 'scene', 'both']) parser.add_argument('--loss_reg_reduction', type=str, default='sum', choices=['mean', 'sum']) parser.add_argument('--loss_reg_squared', default=False, action='store_true') parser.add_argument('--loss_velocity', default=False, action='store_true') # Optimisation parser.add_argument('--init_lr', type=float, default=None) parser.add_argument('--weight_decay', type=float, default=None) # Graph-module-specific flags parser.add_argument('--graph_gnn_layers', type=int, default=2) parser.add_argument('--graph_dropout', type=float, default=0.1) parser.add_argument('--top_n_neighbors', type=int, default=5) parser.add_argument('--rel_traj_hidden', type=int, default=32) parser.add_argument('--y0_score_dim', type=int, default=32, help='Dim of query/key embedding space for RAG scorer.') parser.add_argument('--uncertainty_weight', type=float, default=0.01) return parser.parse_args() # --------------------------------------------------------------------------- # Init # --------------------------------------------------------------------------- def init_basics(args): cfg = Config(args.cfg, f'{args.exp}') tag = '_GRAPHv16' if cfg.denoising_method == 'fm': cfg.sampling_steps = args.sampling_steps if args.fm_skewed_t is not None: cfg.t_schedule = args.fm_skewed_t else: cfg.t_schedule = args.t_schedule if args.t_schedule == 'logit_normal': cfg.logit_norm_mean = args.logit_norm_mean cfg.logit_norm_std = args.logit_norm_std cfg.fm_wrapper = args.fm_wrapper cfg.fm_rew_sqrt = args.fm_rew_sqrt cfg.fm_in_scaling = args.fm_in_scaling if args.fm_skewed_t is not None: tag += f'_FM_S{cfg.sampling_steps}_{cfg.t_schedule}_{cfg.fm_wrapper[:4]}' elif args.t_schedule == 'logit_normal': tag += (f'_FM_S{cfg.sampling_steps}_lnorm' f'_m{cfg.logit_norm_mean}_s{cfg.logit_norm_std}' f'_{cfg.fm_wrapper[:4]}') else: tag += f'_FM_S{cfg.sampling_steps}_uni_{cfg.fm_wrapper[:4]}' if args.drop_method is not None: cfg.drop_method = args.drop_method cfg.drop_logi_k = args.drop_logi_k cfg.drop_logi_m = args.drop_logi_m tag += f'_drop_{cfg.drop_method}_m{cfg.drop_logi_m}_k{cfg.drop_logi_k}' if cfg.fm_rew_sqrt: tag += '_RESQ' if cfg.fm_in_scaling: tag += '_IS' cfg.MODEL.USE_PRE_NORM = args.use_pre_norm cfg.tied_noise = args.tied_noise if args.tied_noise: tag += '_TN' cfg.LOSS_NN_MODE = args.loss_nn_mode cfg.LOSS_REG_REDUCTION = args.loss_reg_reduction cfg.LOSS_REG_SQUARED = args.loss_reg_squared cfg.LOSS_VELOCITY = args.loss_velocity tag += f'_NN_{cfg.LOSS_NN_MODE[:1].upper()}' tag += f'_REG_{cfg.LOSS_REG_REDUCTION[:1].upper()}' if args.loss_reg_squared: tag += '_SQ' if args.loss_velocity: tag += '_VEL' cfg.MODEL.REGRESSION_MLPS[-1] += cfg.MODEL.MODEL_OUT_DIM if args.overfit: tag += '_overfit' if args.n_train != 32500: tag += f'_subset{args.n_train}' cfg.data_norm = args.data_norm tag += f'_{args.data_norm}' if args.init_lr is not None: cfg.OPTIMIZATION.LR = args.init_lr if args.weight_decay is not None: cfg.OPTIMIZATION.WEIGHT_DECAY = args.weight_decay tag += f'_LR{cfg.OPTIMIZATION.LR}_WD{cfg.OPTIMIZATION.WEIGHT_DECAY}' if args.epochs is not None: cfg.OPTIMIZATION.NUM_EPOCHS = args.epochs if args.batch_size is not None: cfg.train_batch_size = args.batch_size cfg.test_batch_size = args.batch_size * 2 if args.checkpt_freq is not None: cfg.checkpt_freq = args.checkpt_freq cfg.max_num_ckpts = args.max_num_ckpts tag += f'_BS{cfg.train_batch_size}_EP{cfg.OPTIMIZATION.NUM_EPOCHS}' cfg.graph_gnn_layers = args.graph_gnn_layers cfg.graph_dropout = args.graph_dropout cfg.top_n_neighbors = args.top_n_neighbors cfg.rel_traj_hidden = args.rel_traj_hidden cfg.y0_score_dim = args.y0_score_dim cfg.uncertainty_weight = args.uncertainty_weight tag += f'_GNN{args.graph_gnn_layers}_N{args.top_n_neighbors}' tag += f'_RTH{args.rel_traj_hidden}_Y0D{args.y0_score_dim}' if args.uncertainty_weight > 0.0: tag += f'_UW{args.uncertainty_weight}' tag = tag.replace('__', '_') cfg.device = 'cuda' if torch.cuda.is_available() else 'cpu' logger = cfg.create_dirs(tag_suffix=tag) if args.fix_random_seed: set_random_seed(args.seed) tb_dir = os.path.abspath(os.path.join(cfg.log_dir, '../tb')) os.makedirs(tb_dir, exist_ok=True) tb_log = SummaryWriter(log_dir=tb_dir) back_up_code_git(cfg, logger=logger) log_config_to_file(cfg.yml_dict, logger=logger) return cfg, logger, tb_log # --------------------------------------------------------------------------- # Data loader # --------------------------------------------------------------------------- def build_data_loader(cfg, args): train_dset = NBADatasetMinMax( data_dir = args.data_dir, obs_len = cfg.past_frames, pred_len = cfg.future_frames, training = True, num_scenes = args.n_train, overfit = args.overfit, cfg = cfg, rotate = args.rotate, data_norm = args.data_norm, ) train_loader = DataLoader( train_dset, batch_size = cfg.train_batch_size, shuffle = True, num_workers = 4, collate_fn = seq_collate_nba_graph, pin_memory = True, ) if args.overfit: test_dset = copy.deepcopy(train_dset) else: test_dset = NBADatasetMinMax( data_dir = args.data_dir, obs_len = cfg.past_frames, pred_len = cfg.future_frames, training = False, overfit = args.overfit, test_scenes = args.n_test, cfg = cfg, rotate = args.rotate, data_norm = args.data_norm, ) test_loader = DataLoader( test_dset, batch_size = cfg.test_batch_size, shuffle = False, num_workers = 4, collate_fn = seq_collate_nba_graph, pin_memory = True, ) return train_loader, test_loader # --------------------------------------------------------------------------- # Network builder # --------------------------------------------------------------------------- def build_network(cfg, args, logger): model = MotionTransformerGraphV16( model_config = cfg.MODEL, logger = logger, config = cfg, graph_num_gnn_layers = args.graph_gnn_layers, graph_dropout = args.graph_dropout, top_n_neighbors = args.top_n_neighbors, rel_traj_hidden = args.rel_traj_hidden, y0_score_dim = args.y0_score_dim, ) if cfg.denoising_method == 'fm': denoiser = FlowMatcher(cfg, model, logger=logger) else: raise NotImplementedError( f'Denoising method [{cfg.denoising_method}] is not implemented.' ) return denoiser # --------------------------------------------------------------------------- # Main # --------------------------------------------------------------------------- def main(): args = parse_config() cfg, logger, tb_log = init_basics(args) train_loader, test_loader = build_data_loader(cfg, args) denoiser = build_network(cfg, args, logger) trainer = Trainer( cfg, denoiser, train_loader, test_loader, tb_log = tb_log, logger = logger, gradient_accumulate_every = 1, ema_decay = 0.995, ema_update_every = 1, ) trainer.train() if __name__ == '__main__': main()