htrnguyen/golf-tech-analysis
0
1from model import EventDetector2import torch3from torch.utils.data import DataLoader4from torchvision import transforms5from dataloader import GolfDB, ToTensor, Normalize6import torch.nn.functional as F7import numpy as np8from util import correct_preds9 10 11def eval(model, split, seq_length, n_cpu, disp):12 dataset = GolfDB(data_file='data/val_split_{}.pkl'.format(split),13 vid_dir='data/videos_160/',14 seq_length=seq_length,15 transform=transforms.Compose([ToTensor(),16 Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])]),17 train=False)18 19 data_loader = DataLoader(dataset,20 batch_size=1,21 shuffle=False,22 num_workers=n_cpu,23 drop_last=False)24 25 correct = []26 27 for i, sample in enumerate(data_loader):28 images, labels = sample['images'], sample['labels']29 # full samples do not fit into GPU memory so evaluate sample in 'seq_length' batches30 batch = 031 while batch * seq_length < images.shape[1]:32 if (batch + 1) * seq_length > images.shape[1]:33 image_batch = images[:, batch * seq_length:, :, :, :]34 else:35 image_batch = images[:, batch * seq_length:(batch + 1) * seq_length, :, :, :]36 logits = model(image_batch.cuda())37 if batch == 0:38 probs = F.softmax(logits.data, dim=1).cpu().numpy()39 else:40 probs = np.append(probs, F.softmax(logits.data, dim=1).cpu().numpy(), 0)41 batch += 142 _, _, _, _, c = correct_preds(probs, labels.squeeze())43 if disp:44 print(i, c)45 correct.append(c)46 PCE = np.mean(correct)47 return PCE48 49 50if __name__ == '__main__':51 52 split = 153 seq_length = 6454 n_cpu = 655 56 model = EventDetector(pretrain=True,57 width_mult=1.,58 lstm_layers=1,59 lstm_hidden=256,60 bidirectional=True,61 dropout=False)62 63 save_dict = torch.load('models_v1/swingnet_1800.pth.tar')64 model.load_state_dict(save_dict['model_state_dict'])65 model.cuda()66 model.eval()67 PCE = eval(model, split, seq_length, n_cpu, True)68 print('Average PCE: {}'.format(PCE))69 70 71 