| """ |
| main_soccer_mid.py — MID baseline on the soccer dataset. |
| |
| Adapted from main_nba_mid.py with: |
| - NUM_AGENTS = 23 (soccer) |
| - Per-scene normalization (abs channel centered by scene centroid at last obs) |
| - No /= (94/28) court rescale (soccer data is already in field-normalized units) |
| - val.npy used as both val and test |
| - Adjusted batch sizes for 23 agents |
| |
| Usage: |
| python main_soccer_mid.py --gpu 0 |
| """ |
|
|
| import os |
| import sys |
| import time |
| import logging |
| import argparse |
| import numpy as np |
| import torch |
| import torch.nn as nn |
| from torch import optim |
| from torch.utils.data import Dataset, DataLoader |
| try: |
| from tensorboardX import SummaryWriter |
| except Exception: |
| from torch.utils.tensorboard import SummaryWriter |
| from tqdm.auto import tqdm |
|
|
| from models.diffusion import DiffusionTraj, VarianceSchedule, TransformerConcatLinear |
|
|
|
|
| OBS_LEN = 10 |
| PRED_LEN = 20 |
| NUM_AGENTS = 23 |
| TRAJ_SCALE = 5.0 |
| TRAJ_MEAN = torch.FloatTensor([-0.726, -0.278]) |
| K_EVAL = 20 |
| PER_SCENE_NORM = True |
|
|
|
|
| class SoccerDatasetMID(Dataset): |
| def __init__(self, data_dir: str, split: str = 'train'): |
| super().__init__() |
| path = os.path.join(data_dir, f'{split}.npy') |
| trajs = np.load(path).astype(np.float32) |
| trajs = torch.from_numpy(trajs).permute(0, 2, 1, 3) |
| self.pre = trajs[:, :, :OBS_LEN, :] |
| self.fut = trajs[:, :, OBS_LEN:, :] |
| print(f'[SoccerDatasetMID] {split}: {path} → {trajs.shape}') |
|
|
| def __len__(self): |
| return len(self.pre) |
|
|
| def __getitem__(self, idx): |
| return self.pre[idx], self.fut[idx] |
|
|
|
|
| def collate_fn(batch): |
| pre = torch.stack([b[0] for b in batch]) |
| fut = torch.stack([b[1] for b in batch]) |
| return pre, fut |
|
|
|
|
| def preprocess_batch(pre_motion, fut_motion, device): |
| B, A = pre_motion.shape[:2] |
| pre = pre_motion.reshape(B * A, OBS_LEN, 2) |
| fut = fut_motion.reshape(B * A, PRED_LEN, 2) |
| last_obs = pre[:, -1:, :] |
|
|
| if PER_SCENE_NORM: |
| scene_center = pre_motion[:, :, -1, :].mean(dim=1, keepdim=True) |
| scene_center = scene_center.unsqueeze(2) |
| abs_xy = ((pre_motion - scene_center) / TRAJ_SCALE).reshape(B * A, OBS_LEN, 2) |
| else: |
| traj_mean = TRAJ_MEAN.to(device) |
| abs_xy = (pre - traj_mean) / TRAJ_SCALE |
|
|
| rel_xy = (pre - last_obs) / TRAJ_SCALE |
| vel_xy = torch.cat([rel_xy[:, 1:] - rel_xy[:, :-1], |
| torch.zeros_like(rel_xy[:, :1])], dim=1) |
|
|
| past_6ch = torch.cat([abs_xy, rel_xy, vel_xy], dim=-1) |
| fut_rel = (fut - last_obs) / TRAJ_SCALE |
|
|
| mask = torch.full((B * A, B * A), float('-inf'), device=device) |
| for i in range(B): |
| s, e = i * A, (i + 1) * A |
| mask[s:e, s:e] = 0.0 |
|
|
| return past_6ch, fut_rel, mask, last_obs |
|
|
|
|
| class _STEncoder(nn.Module): |
| def __init__(self, in_channels=6, hidden=256): |
| super().__init__() |
| self.conv = nn.Conv1d(in_channels, 32, kernel_size=3, stride=1, padding=1) |
| self.relu = nn.ReLU() |
| self.gru = nn.GRU(32, hidden, num_layers=1, batch_first=True) |
| nn.init.kaiming_normal_(self.conv.weight) |
| nn.init.kaiming_normal_(self.gru.weight_ih_l0) |
| nn.init.kaiming_normal_(self.gru.weight_hh_l0) |
| nn.init.zeros_(self.conv.bias) |
| nn.init.zeros_(self.gru.bias_ih_l0) |
| nn.init.zeros_(self.gru.bias_hh_l0) |
|
|
| def forward(self, x): |
| h = self.relu(self.conv(x.transpose(1, 2))) |
| _, state = self.gru(h.transpose(1, 2)) |
| return state.squeeze(0) |
|
|
|
|
| class _SocialTransformer(nn.Module): |
| def __init__(self, past_len=OBS_LEN, hidden=256): |
| super().__init__() |
| self.proj = nn.Linear(past_len * 6, hidden, bias=False) |
| layer = nn.TransformerEncoderLayer( |
| d_model=hidden, nhead=2, |
| dim_feedforward=hidden, batch_first=False) |
| self.encoder = nn.TransformerEncoder(layer, num_layers=2) |
|
|
| def forward(self, x_flat, mask): |
| h = self.proj(x_flat).unsqueeze(1) |
| h = h + self.encoder(h, mask) |
| return h.squeeze(1) |
|
|
|
|
| class SoccerEncoder(nn.Module): |
| def __init__(self, encoder_dim=256, past_len=OBS_LEN): |
| super().__init__() |
| self.ego_encoder = _STEncoder(in_channels=6, hidden=256) |
| self.social_encoder = _SocialTransformer(past_len=past_len, hidden=256) |
| self.fusion = nn.Linear(512, encoder_dim) |
|
|
| def forward(self, past_6ch, social_mask): |
| ego = self.ego_encoder(past_6ch) |
| social = self.social_encoder( |
| past_6ch.reshape(past_6ch.size(0), -1), social_mask) |
| return self.fusion(torch.cat([ego, social], dim=-1)) |
|
|
|
|
| class Trainer: |
| def __init__(self, args): |
| self.args = args |
| self.device = torch.device(f'cuda:{args.gpu}' if torch.cuda.is_available() else 'cpu') |
| self._build_dirs() |
| self._build_data() |
| self._build_model() |
| self._build_optimizer() |
|
|
| def _build_dirs(self): |
| self.exp_dir = os.path.join('experiments', self.args.exp_name) |
| os.makedirs(self.exp_dir, exist_ok=True) |
| self.tb_log = SummaryWriter(log_dir=self.exp_dir) |
| log_path = os.path.join( |
| self.exp_dir, |
| 'soccer_{}.log'.format(time.strftime('%Y-%m-%d-%H-%M'))) |
| self.log = logging.getLogger(self.args.exp_name) |
| self.log.setLevel(logging.INFO) |
| self.log.addHandler(logging.FileHandler(log_path)) |
| self.log.addHandler(logging.StreamHandler(sys.stdout)) |
| self.log.info(f"Args: {self.args}") |
|
|
| def _build_data(self): |
| train_dset = SoccerDatasetMID(self.args.data_dir, split='train') |
| test_dset = SoccerDatasetMID(self.args.data_dir, split='val') |
| self.train_loader = DataLoader( |
| train_dset, batch_size=self.args.batch_size, |
| shuffle=True, num_workers=4, |
| collate_fn=collate_fn, pin_memory=True) |
| self.test_loader = DataLoader( |
| test_dset, batch_size=self.args.eval_batch_size, |
| shuffle=False, num_workers=4, |
| collate_fn=collate_fn, pin_memory=True) |
| self.log.info(f"Train: {len(train_dset)} Val/Test: {len(test_dset)}") |
|
|
| def _build_model(self): |
| self.encoder = SoccerEncoder( |
| encoder_dim=self.args.encoder_dim, past_len=OBS_LEN, |
| ).to(self.device) |
| net = TransformerConcatLinear( |
| point_dim=2, context_dim=self.args.encoder_dim, |
| tf_layer=self.args.tf_layer, residual=False) |
| self.diffusion = DiffusionTraj( |
| net=net, |
| var_sched=VarianceSchedule(num_steps=100, beta_T=5e-2, mode='linear'), |
| ).to(self.device) |
| n_enc = sum(p.numel() for p in self.encoder.parameters()) |
| n_diff = sum(p.numel() for p in self.diffusion.parameters()) |
| self.log.info(f"Encoder: {n_enc:,} Diffusion: {n_diff:,}") |
|
|
| def _build_optimizer(self): |
| params = list(self.encoder.parameters()) + list(self.diffusion.parameters()) |
| self.optimizer = optim.Adam(params, lr=self.args.lr) |
| self.scheduler = optim.lr_scheduler.ExponentialLR(self.optimizer, gamma=0.98) |
|
|
| def train(self): |
| best_ade = float('inf') |
| for epoch in range(1, self.args.epochs + 1): |
| self.encoder.train() |
| self.diffusion.train() |
| total_loss, count = 0.0, 0 |
| pbar = tqdm(self.train_loader, ncols=90) |
| for pre, fut in pbar: |
| pre, fut = pre.to(self.device), fut.to(self.device) |
| past_6ch, fut_rel, mask, _ = preprocess_batch(pre, fut, self.device) |
| context = self.encoder(past_6ch, mask) |
| loss = self.diffusion.get_loss(fut_rel, context) |
| self.optimizer.zero_grad() |
| loss.backward() |
| nn.utils.clip_grad_norm_( |
| list(self.encoder.parameters()) + |
| list(self.diffusion.parameters()), 1.0) |
| self.optimizer.step() |
| total_loss += loss.item() |
| count += 1 |
| pbar.set_description(f"Epoch {epoch} loss={total_loss/count:.4f}") |
|
|
| self.scheduler.step() |
| avg_loss = total_loss / count |
| self.tb_log.add_scalar('loss/train', avg_loss, epoch) |
| self.log.info(f"Epoch {epoch} train_loss={avg_loss:.4f}") |
|
|
| if epoch % self.args.eval_every == 0: |
| ade, fde = self.evaluate() |
| self.tb_log.add_scalar('metric/ADE', ade, epoch) |
| self.tb_log.add_scalar('metric/FDE', fde, epoch) |
| self.log.info(f"Epoch {epoch} ADE={ade:.4f} FDE={fde:.4f}") |
| if ade < best_ade: |
| best_ade = ade |
| torch.save({ |
| 'encoder': self.encoder.state_dict(), |
| 'diffusion': self.diffusion.state_dict(), |
| 'epoch': epoch, 'ade': ade, 'fde': fde, |
| }, os.path.join(self.exp_dir, 'best.pt')) |
| self.log.info(f" ** New best ADE={ade:.4f} FDE={fde:.4f}") |
|
|
| @torch.no_grad() |
| def evaluate(self): |
| self.encoder.eval() |
| self.diffusion.eval() |
| ade_sum, fde_sum, n_agents = 0.0, 0.0, 0 |
| for pre, fut in tqdm(self.test_loader, ncols=90, desc='Eval'): |
| pre, fut = pre.to(self.device), fut.to(self.device) |
| B = pre.size(0) |
| past_6ch, _, mask, last_obs = preprocess_batch(pre, fut, self.device) |
| context = self.encoder(past_6ch, mask) |
| pred_rel = self.diffusion.sample( |
| num_points=PRED_LEN, context=context, |
| sample=K_EVAL, bestof=True, |
| sampling=self.args.sampling, step=self.args.sampling_step) |
| pred_abs = pred_rel * TRAJ_SCALE + last_obs.unsqueeze(0) |
| fut_abs = fut.reshape(B * NUM_AGENTS, PRED_LEN, 2) |
| dist = (pred_abs - fut_abs.unsqueeze(0)).norm(dim=-1) |
| ade_sum += dist.mean(dim=-1).min(dim=0).values.sum().item() |
| fde_sum += dist[:, :, -1].min(dim=0).values.sum().item() |
| n_agents += B * NUM_AGENTS |
| return ade_sum / n_agents, fde_sum / n_agents |
|
|
|
|
| def parse_args(): |
| p = argparse.ArgumentParser() |
| p.add_argument('--data_dir', type=str, |
| default='/mnt/jaewoo4tb/srtp/srtp/raw_data/soccer') |
| p.add_argument('--exp_name', type=str, default='mid_soccer_baseline') |
| p.add_argument('--gpu', type=int, default=0) |
| p.add_argument('--epochs', type=int, default=100) |
| p.add_argument('--batch_size', type=int, default=32) |
| p.add_argument('--eval_batch_size', type=int, default=64) |
| p.add_argument('--lr', type=float, default=1e-3) |
| p.add_argument('--eval_every', type=int, default=5) |
| p.add_argument('--encoder_dim', type=int, default=256) |
| p.add_argument('--tf_layer', type=int, default=3) |
| p.add_argument('--sampling', type=str, default='ddim', choices=['ddpm', 'ddim']) |
| p.add_argument('--sampling_step', type=int, default=10) |
| return p.parse_args() |
|
|
|
|
| if __name__ == '__main__': |
| args = parse_args() |
| trainer = Trainer(args) |
| trainer.train() |
|
|