(model_name, dataset_name, attack_name, defense_name, model, model_path, benign_dataset, attacked_dataset, defense, y_target)
| 31 | |
| 32 | |
| 33 | def 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 ========== |
no test coverage detected