-
Notifications
You must be signed in to change notification settings - Fork 7
Expand file tree
/
Copy pathvisualization.py
More file actions
100 lines (80 loc) · 2.84 KB
/
Copy pathvisualization.py
File metadata and controls
100 lines (80 loc) · 2.84 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
88
89
90
91
92
93
94
95
96
97
98
99
100
""""
Author: Xu Ma
Date: Aug/15/2019
Email: xuma@my.unt.edu
Useage:
"""
import argparse
import os
import shutil
import time
import math
import traceback
import copy
import os.path as osp
import click
import cv2
import matplotlib.cm as cm
import numpy as np
import torch
import torch.hub
import torch.nn.functional as F
from torch.autograd import Variable
from torchvision import transforms
import sys
import torch
from torch.autograd import Variable
import torch.nn as nn
import torch.nn.parallel
import torch.backends.cudnn as cudnn
import torch.distributed as dist
import torch.optim
import torch.utils.data
import torch.utils.data.distributed
import models as models
from utils import Logger, mkdir_p, get_device, get_classtable,GradCAM
try:
from nvidia.dali.plugin.pytorch import DALIClassificationIterator
from nvidia.dali.pipeline import Pipeline
import nvidia.dali.ops as ops
import nvidia.dali.types as types
except ImportError:
raise ImportError("Please install DALI from https://www.github.com/NVIDIA/DALI to run this example.")
try:
from apex.parallel import DistributedDataParallel as DDP
from apex.fp16_utils import *
from apex import amp, optimizers
except ImportError:
raise ImportError("Please install apex from https://www.github.com/nvidia/apex to run this example.")
model_names = sorted(name for name in models.__dict__
if name.islower() and not name.startswith("__")
and callable(models.__dict__[name]))
parser = argparse.ArgumentParser(description='PyTorch ImageNet Training')
parser.add_argument('-d', '--data', default='/PATH_to_imageNet/ImageNet2012/', type=str)
parser.add_argument('--arch', '-a', metavar='ARCH', default='resnet18', choices=model_names)
parser.add_argument('-c', '--checkpoint', type=str, metavar='PATH',
help='path to your checkpoint')
parser.add_argument('--cuda', action='store_true', default=True,
help='Running on GPU or CPU.')
parser.add_argument('-i', '--image', default='cat.png', type=str, metavar='PATH',
help='path to your image')
parser.add_argument('-o', '--output-dir', default='', type=str, metavar='PATH',
help='folder for save images')
parser.add_argument('-t', '--target-layer', default='layer4', type=str,
help='Target layer for visualization')
args = parser.parse_args()
def main():
device = get_device(args.cuda)
# Synset words
classes = get_classtable()
# Model from torchvision
model = models.__dict__[args.arch]()
model = model.cuda()
check_point = torch.load(args.checkpoint, map_location=lambda storage, loc: storage.cuda(0))
model.load_state_dict(check_point['state_dict'])
model.to(device)
model.eval()
gcam = GradCAM(model=model)
_ = gcam.forward(images)
if __name__ == '__main__':
main()