htrnguyen/golf-tech-analysis
0
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 