File size: 1,575 Bytes
0d52562 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 | """SDD dataloader for LED. Variable-A scenes, batch_size=1."""
import pickle, torch
from torch.utils.data import Dataset
class SDDDataset(Dataset):
"""Each sample is a scene with A=1+N agents (target + neighbors), shape [A,20,2]."""
def __init__(self, obs_len=8, pred_len=12, split='train'):
super().__init__()
self.obs_len, self.pred_len = obs_len, pred_len
path = f'data/files/sdd_{split}_v2.pkl'
with open(path, 'rb') as f:
self.scenes = pickle.load(f)
a = [s.shape[0] for s in self.scenes]
print(f'[SDDDataset] {split}: {len(self.scenes)} scenes, '
f'A min/mean/max = {min(a)}/{sum(a)/len(a):.1f}/{max(a)}')
def __len__(self): return len(self.scenes)
def __getitem__(self, idx):
traj = torch.from_numpy(self.scenes[idx]) # [A, 20, 2]
pre = traj[:, :self.obs_len] # [A, 8, 2]
fut = traj[:, self.obs_len:] # [A, 12, 2]
A = traj.size(0)
pre_mask = torch.ones(A, self.obs_len)
fut_mask = torch.ones(A, self.pred_len)
return [pre, fut, pre_mask, fut_mask]
def sdd_seq_collate(batch):
assert len(batch) == 1, 'SDD uses batch_size=1 (variable A)'
pre, fut, pm, fm = batch[0]
return {
'pre_motion_3D': pre.unsqueeze(0), # [1, A, 8, 2]
'fut_motion_3D': fut.unsqueeze(0), # [1, A, 12, 2]
'pre_motion_mask': pm.unsqueeze(0),
'fut_motion_mask': fm.unsqueeze(0),
'traj_scale': 1,
'pred_mask': None,
'seq': 'sdd',
}
|