| """ |
| 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_v9b import MotionTransformerGraphV9b |
| from trainer.denoising_model_trainers import Trainer |
|
|
|
|
| |
| |
| |
|
|
| def parse_config(): |
| parser = argparse.ArgumentParser() |
|
|
| |
| parser.add_argument('--cfg', default='cfg/nba/cor_fm.yml', type=str) |
| parser.add_argument('--exp', default='', type=str) |
|
|
| |
| 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']) |
|
|
| |
| parser.add_argument('--fix_random_seed', action='store_true', default=False) |
| parser.add_argument('--seed', type=int, default=42) |
|
|
| |
| 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') |
|
|
| |
| 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) |
|
|
| |
| parser.add_argument('--use_pre_norm', default=False, action='store_true') |
|
|
| |
| parser.add_argument('--tied_noise', default=False, action='store_true') |
|
|
| |
| 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') |
|
|
| |
| parser.add_argument('--init_lr', type=float, default=None) |
| parser.add_argument('--weight_decay', type=float, default=None) |
|
|
| |
| 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() |
|
|
|
|
| |
| |
| |
|
|
| def init_basics(args): |
| cfg = Config(args.cfg, f'{args.exp}') |
| tag = '_GRAPHv9b' |
|
|
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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 |
|
|
|
|
| |
| |
| |
|
|
| def build_network(cfg, args, logger): |
| model = MotionTransformerGraphV9b( |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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() |
|
|