MCPcopy Create free account
hub / github.com/THUYimingLi/BackdoorBox / test

Function test

tests/test_AutoEncoder.py:33–87  ·  view source on GitHub ↗
(model_name, dataset_name, attack_name, defense_name, model, model_path, benign_dataset, attacked_dataset, defense, y_target)

Source from the content-addressed store, hash-verified

31
32
33def test(model_name, dataset_name, attack_name, defense_name, model, model_path, benign_dataset, attacked_dataset, defense, y_target):
34 if dataset_name == 'CIFAR-10':
35 data = any2tensor(benign_dataset.data)
36 data = data.permute((0, 3, 1, 2))
37 data = defense.preprocess(data.float() / 255)
38
39 schedule = {
40 'device': 'GPU',
41 'CUDA_VISIBLE_DEVICES': CUDA_VISIBLE_DEVICES,
42 'GPU_num': 1,
43
44 'test_model': model_path,
45 'batch_size': batch_size,
46 'num_workers': num_workers,
47 }
48 res = defense.predict(model, data.float() / 255, schedule)
49
50 schedule = {
51 'device': 'GPU',
52 'CUDA_VISIBLE_DEVICES': CUDA_VISIBLE_DEVICES,
53 'GPU_num': 1,
54
55 'test_model': model_path,
56 'batch_size': batch_size,
57 'num_workers': num_workers,
58
59 'metric': 'BA',
60
61 'save_dir': 'experiments',
62 'experiment_name': f'{model_name}_{dataset_name}_{attack_name}_{defense_name}_BA'
63 }
64 defense.test(model, benign_dataset, schedule)
65
66 schedule = {
67 'device': 'GPU',
68 'CUDA_VISIBLE_DEVICES': CUDA_VISIBLE_DEVICES,
69 'GPU_num': 1,
70
71 'test_model': model_path,
72 'batch_size': batch_size,
73 'num_workers': num_workers,
74
75 # 1. ASR: the attack success rate calculated on all poisoned samples
76 # 2. ASR_NoTarget: the attack success rate calculated on all poisoned samples whose ground-truth labels are not the target label
77 # 3. BA: the accuracy on all benign samples
78 # Hint: For ASR and BA, the computation of the metric is decided by the dataset but not schedule['metric'].
79 # In other words, ASR or BA does not influence the computation of the metric.
80 # For ASR_NoTarget, the code will delete all the samples whose ground-truth labels are the target label and then compute the metric.
81 'metric': 'ASR_NoTarget',
82 'y_target': y_target,
83
84 'save_dir': 'experiments',
85 'experiment_name': f'{model_name}_{dataset_name}_{attack_name}_{defense_name}_ASR'
86 }
87 defense.test(model, attacked_dataset, schedule)
88
89
90# ========== ResNet-18_CIFAR-10_Attack_AutoEncoder ==========

Callers 1

Calls 4

any2tensorFunction · 0.90
preprocessMethod · 0.45
predictMethod · 0.45
testMethod · 0.45

Tested by

no test coverage detected