CoolFace
Apppublic

htrnguyen/golf-tech-analysis

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
eval.py71 linesDownload Raw Back to src
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