CoolFace
Apppublic

htrnguyen/golf-tech-analysis

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
dataloader.py114 linesDownload Raw Back to src
1import os.path as osp2import cv23import pandas as pd4import numpy as np5import torch6from torch.utils.data import Dataset, DataLoader7from torchvision import transforms8 9 10class GolfDB(Dataset):11    def __init__(self, data_file, vid_dir, seq_length, transform=None, train=True):12        self.df = pd.read_pickle(data_file)13        self.vid_dir = vid_dir14        self.seq_length = seq_length15        self.transform = transform16        self.train = train17 18    def __len__(self):19        return len(self.df)20 21    def __getitem__(self, idx):22        a = self.df.loc[idx, :]  # annotation info23        events = a['events']24        events -= events[0]  # now frame #s correspond to frames in preprocessed video clips25 26        images, labels = [], []27        cap = cv2.VideoCapture(osp.join(self.vid_dir, '{}.mp4'.format(a['id'])))28 29        if self.train:30            # random starting position, sample 'seq_length' frames31            start_frame = np.random.randint(events[-1] + 1)32            cap.set(cv2.CAP_PROP_POS_FRAMES, start_frame)33            pos = start_frame34            while len(images) < self.seq_length:35                ret, img = cap.read()36                if ret:37                    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)38                    images.append(img)39                    if pos in events[1:-1]:40                        labels.append(np.where(events[1:-1] == pos)[0][0])41                    else:42                        labels.append(8)43                    pos += 144                else:45                    cap.set(cv2.CAP_PROP_POS_FRAMES, 0)46                    pos = 047            cap.release()48        else:49            # full clip50            for pos in range(int(cap.get(cv2.CAP_PROP_FRAME_COUNT))):51                _, img = cap.read()52                img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)53                images.append(img)54                if pos in events[1:-1]:55                    labels.append(np.where(events[1:-1] == pos)[0][0])56                else:57                    labels.append(8)58            cap.release()59 60        sample = {'images':np.asarray(images), 'labels':np.asarray(labels)}61        if self.transform:62            sample = self.transform(sample)63        return sample64 65 66class ToTensor(object):67    """Convert ndarrays in sample to Tensors."""68    def __call__(self, sample):69        images, labels = sample['images'], sample['labels']70        images = images.transpose((0, 3, 1, 2))71        return {'images': torch.from_numpy(images).float().div(255.),72                'labels': torch.from_numpy(labels).long()}73 74 75class Normalize(object):76    def __init__(self, mean, std):77        self.mean = torch.tensor(mean, dtype=torch.float32)78        self.std = torch.tensor(std, dtype=torch.float32)79 80    def __call__(self, sample):81        images, labels = sample['images'], sample['labels']82        images.sub_(self.mean[None, :, None, None]).div_(self.std[None, :, None, None])83        return {'images': images, 'labels': labels}84 85 86if __name__ == '__main__':87 88    norm = Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])  # ImageNet mean and std (RGB)89 90    dataset = GolfDB(data_file='data/train_split_1.pkl',91                     vid_dir='data/videos_160/',92                     seq_length=64,93                     transform=transforms.Compose([ToTensor(), norm]),94                     train=False)95 96    data_loader = DataLoader(dataset, batch_size=1, shuffle=False, num_workers=6, drop_last=False)97 98    for i, sample in enumerate(data_loader):99        images, labels = sample['images'], sample['labels']100        events = np.where(labels.squeeze() < 8)[0]101        print('{} events: {}'.format(len(events), events))102 103 104 105 106    107 108 109 110 111 112       113 114