sra-trajectory-code / MID /main_soccer_mid.py
po03087's picture
SRA: MID/LED/MoFlow code + RUNNING.md instructions (code only, no data/ckpts)
d4cbafd verified
Raw
History Blame Contribute Delete
11.3 kB
"""
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) # (N, 30, 23, 2)
trajs = torch.from_numpy(trajs).permute(0, 2, 1, 3) # (N, 23, 30, 2)
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) # [B, 1, 2]
scene_center = scene_center.unsqueeze(2) # [B, 1, 1, 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()