sra-trajectory-code / MoFlow /data /dataloader_sport.py
po03087's picture
FIX: restore data/ dataloader modules (lost by over-broad --exclude=data)
0d52562 verified
Raw
History Blame Contribute Delete
5.56 kB
"""
Sport-dataset dataloader for MoFlow (soccer / football).
Mirrors the structure of dataloader_nba.py's NBADatasetMinMax but:
- Loads from absolute sport data_dir (train.npy / val.npy)
- No 94/28 court rescale (soccer/football data is already in field units)
- Agent count taken from cfg.agents (23 for soccer/football)
- Supports per-scene normalization of the abs channel (matching the LED/MID sport setup)
- Soccer has no test split → we use val.npy as both val and test
"""
import os
import numpy as np
import torch
from torch.utils.data import Dataset
from utils.normalization import normalize_min_max, normalize_sqrt
from utils.utils import rotate_trajs_x_direction
def seq_collate_sport(batch):
(past_traj, fut_traj, past_traj_orig, fut_traj_orig, traj_vel) = zip(*batch)
pre_motion_3D = torch.stack(past_traj, dim=0)
fut_motion_3D = torch.stack(fut_traj, dim=0)
pre_motion_3D_orig = torch.stack(past_traj_orig, dim=0)
fut_motion_3D_orig = torch.stack(fut_traj_orig, dim=0)
fut_traj_vel = torch.stack(traj_vel, dim=0)
B = pre_motion_3D.shape[0]
A = pre_motion_3D.shape[1] # per-sample agent count
batch_size = torch.tensor(B)
traj_mask = torch.zeros(B * A, B * A)
for i in range(B):
traj_mask[i * A:(i + 1) * A, i * A:(i + 1) * A] = 1.
return {
'batch_size': batch_size,
'past_traj': pre_motion_3D,
'fut_traj': fut_motion_3D,
'past_traj_original_scale': pre_motion_3D_orig,
'fut_traj_original_scale': fut_motion_3D_orig,
'traj_mask': traj_mask,
'fut_traj_vel': fut_traj_vel,
}
class SportDatasetMinMax(Dataset):
def __init__(
self,
obs_len=10,
pred_len=20,
training=True,
num_scenes=None,
test_scenes=None,
overfit=False,
cfg=None,
data_dir='',
rotate=False,
data_norm='min_max',
):
super().__init__()
self.obs_len = obs_len
self.pred_len = pred_len
self.seq_len = obs_len + pred_len
self.traj_mean = torch.FloatTensor(cfg.traj_mean).unsqueeze(0).unsqueeze(0).unsqueeze(0)
self.per_scene_norm = bool(cfg.get('per_scene_norm', False))
# Soccer has no test split → use val.npy for both roles.
split = 'train' if training else 'val'
path = os.path.join(data_dir, f'{split}.npy')
if not os.path.isfile(path):
raise FileNotFoundError(path)
trajs = np.load(path).astype(np.float32) # (N, 30, A, 2)
if num_scenes is not None and training:
trajs = trajs[:num_scenes]
if test_scenes is not None and not training:
trajs = trajs[:test_scenes]
self.data_len = len(trajs)
print(f'[SportDatasetMinMax] {split}: {path}{trajs.shape}')
# (N, A, T, 2)
self.traj_abs = torch.from_numpy(trajs).type(torch.float).permute(0, 2, 1, 3)
self.actor_num = self.traj_abs.shape[1]
pre_motion_3D = self.traj_abs[:, :, :self.obs_len, :]
fut_motion_3D = self.traj_abs[:, :, self.obs_len:, :]
initial_pos = pre_motion_3D[:, :, -1:]
fut_traj = (fut_motion_3D - initial_pos).contiguous()
if self.per_scene_norm:
# Scene centroid at the last observed frame: [N, 1, 1, 2]
scene_center = pre_motion_3D[:, :, -1, :].mean(dim=1, keepdim=True).unsqueeze(2)
past_traj_abs = (pre_motion_3D - scene_center).contiguous()
else:
past_traj_abs = (pre_motion_3D - self.traj_mean).contiguous()
past_traj_rel = (pre_motion_3D - initial_pos).contiguous()
if rotate:
past_traj_rel, fut_traj, past_traj_abs = rotate_trajs_x_direction(
past_traj_rel, fut_traj, past_traj_abs)
past_traj_vel = torch.cat(
(past_traj_rel[:, :, 1:] - past_traj_rel[:, :, :-1],
torch.zeros_like(past_traj_rel[:, :, -1:])), dim=2)
past_traj = torch.cat((past_traj_abs, past_traj_rel, past_traj_vel), dim=-1)
self.fut_traj_vel = torch.cat(
(fut_traj[:, :, 1:] - fut_traj[:, :, :-1],
torch.zeros_like(fut_traj[:, :, -1:])), dim=2)
if training:
cfg.fut_traj_max = fut_traj.max()
cfg.fut_traj_min = fut_traj.min()
cfg.past_traj_max = past_traj.max()
cfg.past_traj_min = past_traj.min()
self.past_traj_original_scale = past_traj
self.fut_traj_original_scale = fut_traj
self.data_norm = data_norm
if data_norm == 'min_max':
self.past_traj = normalize_min_max(
past_traj, cfg.past_traj_min, cfg.past_traj_max, -1, 1).contiguous()
self.fut_traj = normalize_min_max(
fut_traj, cfg.fut_traj_min, cfg.fut_traj_max, -1, 1).contiguous()
elif data_norm == 'sqrt':
sqrt_a_ = torch.tensor([cfg.sqrt_x_a, cfg.sqrt_y_a], device=past_traj.device)
sqrt_b_ = torch.tensor([cfg.sqrt_x_b, cfg.sqrt_y_b], device=past_traj.device)
self.past_traj = past_traj
self.fut_traj = normalize_sqrt(fut_traj, sqrt_a_, sqrt_b_).contiguous()
def __len__(self):
return self.data_len
def __getitem__(self, index):
return [
self.past_traj[index],
self.fut_traj[index],
self.past_traj_original_scale[index],
self.fut_traj_original_scale[index],
self.fut_traj_vel[index],
]