"""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', }