-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest_aa.py
More file actions
executable file
·68 lines (57 loc) · 2.67 KB
/
Copy pathtest_aa.py
File metadata and controls
executable file
·68 lines (57 loc) · 2.67 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
import torch
import torchvision
import os
os.environ["CUDA_VISIBLE_DEVICES"] = "0"
from models.resnet import ResNet18
from models.wideresnet import WideResNet
import argparse
parser = argparse.ArgumentParser()
parser.add_argument('--model-path', default=None, type=str, )
parser.add_argument('--log', default=None, type=str, )
parser.add_argument('--eps', default=8/255, type=float, )
parser.add_argument('--dataset', default="cifar10", type=str, )
parser.add_argument('--arch', default="resnet18", type=str, )
parser.add_argument('--attacks', type=str, default='ensemble1')
args = parser.parse_args()
if args.attacks=='ensemble1':
attacks = ['apgd-ce', 'apgd-t', 'fab-t', 'square']
elif args.attacks=='ensemble2':
attacks = ['apgd-ce', 'apgd-dlr', 'fab', 'square']
else:
raise "invalid attack ensemble"
print("epsilon:", args.eps)
print("attacks to run: ", args.attacks)
if args.dataset == "cifar10":
testset = torchvision.datasets.CIFAR10("~/datasets/", train=False, transform=torchvision.transforms.ToTensor())
elif args.dataset == "cifar100":
testset = torchvision.datasets.CIFAR100("~/datasets/", train=False, transform=torchvision.transforms.ToTensor())
elif args.dataset == "svhn":
testset = torchvision.datasets.SVHN(root='~/datasets/', split="test", download=False, transform=torchvision.transforms.ToTensor())
testloader = torch.utils.data.DataLoader(testset, shuffle=False, batch_size=len(testset),
num_workers=8, pin_memory=True)
num_classes = 10 if args.dataset != "cifar100" else 100
if args.arch == "resnet18":
model = ResNet18(num_classes=num_classes)
model.cuda()
#model = torch.nn.DataParallel(model)
state_dict = torch.load(args.model_path)["model_state_dict"]
model.load_state_dict(state_dict)
model.eval()
elif args.arch == "wrn-34-10":
state_dict = torch.load(args.model_path, map_location="cpu")
if "model_state_dict" in list(state_dict.keys()):
state_dict = state_dict["model_state_dict"]
if "sub_block1.layer.0.bn1.weight" in list(state_dict.keys()):
subblock1=True
else:
subblock1=False
model = WideResNet(depth=34, widen_factor=10, subblock1=subblock1, num_classes=num_classes)
model.cuda()
#model = torch.nn.DataParallel(model)
model.load_state_dict(state_dict)
model.eval()
from autoattack import AutoAttack
#adversary = AutoAttack(model, norm='Linf', eps=args.eps, log_path=args.log, version='standard')
adversary = AutoAttack(model, norm='Linf', eps=args.eps, log_path=args.log, version='custom', attacks_to_run=attacks)
for i, (x_test, y_test) in enumerate(testloader):
x_adv = adversary.run_standard_evaluation(x_test, y_test, bs=1000)