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