-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy patheval_imagenet.py
More file actions
87 lines (66 loc) · 3.49 KB
/
Copy patheval_imagenet.py
File metadata and controls
87 lines (66 loc) · 3.49 KB
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
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
import argparse
from pathlib import Path
import numpy as np
import matplotlib.pyplot as plt
import os
import time
from tqdm import tqdm
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader
import torchvision.datasets as datasets
import bagnetsv2 as bagnets # Change to import bagnets to see the results of the original model
import pretrain_imagenet as pt
import utils
from torchvision.models import get_model as tv_get_model
def get_args():
parser = argparse.ArgumentParser(description='Evaluation on ImageNet.')
parser.add_argument('--backbone', default='bagnet33', type=str, help='backbone model', choices=['resnet18', 'resnet50', 'bagnet33', 'bagnet17', 'bagnet9'])
parser.add_argument('--dataset', default='imagenet', type=str, help='dataset to train on', choices=['imagenet', 'imagenette'])
parser.add_argument('--checkpoint', default='checkpoints/bagnet33_imagenet_pretrained.pt', type=str, help='filename of the checkpoint to evaluate')
parser.add_argument('--imagesize', default=224, type=int, help='image size, only square images are supported')
parser.add_argument('--batchsize', default=256, type=int, help='batch size for training')
parser.add_argument('--numworkers', default=4, type=int, help='number of subprocesses to use for dataloading')
parser.add_argument('--device', default='cuda:0', type=str, help='device in which training will take place')
args = parser.parse_args()
args.imagesize = (args.imagesize, args.imagesize)
return args
def get_dataloader(args):
transform = utils.get_augmentations(args.imagesize, normalization=utils.IMAGENET_NORMALIZATION, imagenet=True)['test']
if args.dataset == 'imagenet':
dataset_test = datasets.ImageNet(utils.IMAGENET_DIR, split='val', transform=transform)
n_classes = 1000
else:
dataset_test = datasets.Imagenette('datasets', split='val', transform=transform, size='320px', download=True)
n_classes = 10
dataloader_test = DataLoader(dataset_test, batch_size=args.batchsize, shuffle=False, num_workers=args.numworkers, pin_memory=True, drop_last=False)
return dataloader_test, n_classes
def load_model(args, n_classes):
if 'bagnet' in args.backbone:
model = bagnets.get_bagnet(args.backbone, weights=None, num_classes=n_classes)
else:
model = tv_get_model(args.backbone, weights=None)
model.fc = nn.Linear(model.fc.in_features, n_classes)
checkpoint = torch.load(args.checkpoint, map_location=torch.device('cpu'))
model.load_state_dict(checkpoint['state_dict'])
return model
if __name__ == '__main__':
args = get_args()
start = time.perf_counter()
dataloader_test, n_classes = get_dataloader(args)
model = load_model(args, n_classes)
# Check for dead convolutional layers
dead_layer_count = 0
for name, parameters in model.named_parameters():
if 'conv' in name:
max_weight = parameters.flatten().abs().max()
if max_weight <= 1e-4:
dead_layer_count += 1
print(f'Dead layer count (max(abs(parameters) <= 1e-4 ) = {dead_layer_count}')
# Accuracy
preds, probs, targets = pt.predict(model, dataloader_test, args.device)
acc = utils.accuracy(torch.from_numpy(probs), torch.from_numpy(targets), (1, 5))
print(f'Top 1 accuracy on validation set: {acc[0].item():.2f}')
print(f'Top 5 accuracy on validation set: {acc[1].item():.2f}')
print(f'Total testing time: {time.perf_counter() - start:.1f} s')