sra-trajectory-code / LED /data /dataloader_sdd.py
po03087's picture
FIX: restore data/ dataloader modules (lost by over-broad --exclude=data)
0d52562 verified
Raw
History Blame Contribute Delete
1.58 kB
"""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',
}