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