(
in_size: int,
classes: int,
labels: list,
proportion: float = 0.0,
trigger: str = None,
attack: str = None,
src_label: int = None,
tar_label: int = None)
| 334 | |
| 335 | |
| 336 | def prepare_backdoor_attack( |
| 337 | in_size: int, |
| 338 | classes: int, |
| 339 | labels: list, |
| 340 | proportion: float = 0.0, |
| 341 | trigger: str = None, |
| 342 | attack: str = None, |
| 343 | src_label: int = None, |
| 344 | tar_label: int = None) -> tuple: |
| 345 | |
| 346 | if (trigger is None) or (attack is None): |
| 347 | return None, None, None |
| 348 | |
| 349 | mask = np.ones([in_size, in_size], dtype=np.uint8) |
| 350 | pattern = np.zeros([in_size, in_size], dtype=np.uint8) |
| 351 | |
| 352 | if trigger == 'single-pixel': |
| 353 | mask[-2, -2], pattern[-2, -2] = 0, 255 |
| 354 | |
| 355 | elif trigger == 'pattern': |
| 356 | mask[-2, -2], pattern[-2, -2] = 0, 255 |
| 357 | mask[-2, -4], pattern[-2, -4] = 0, 255 |
| 358 | mask[-4, -2], pattern[-4, -2] = 0, 255 |
| 359 | mask[-3, -3], pattern[-3, -3] = 0, 255 |
| 360 | |
| 361 | else: |
| 362 | raise ValueError( |
| 363 | 'backdoor trigger {} is not supported'.format(trigger)) |
| 364 | |
| 365 | trigger_trans = TriggerTrans(mask, pattern) |
| 366 | |
| 367 | if attack == 'single-target': |
| 368 | indices = np.where(np.array(labels) == src_label)[0] |
| 369 | attack_trans = SingleTargetTrans(src_label, tar_label) |
| 370 | |
| 371 | elif attack == 'all-to-all': |
| 372 | indices = np.arange(len(labels)) |
| 373 | attack_trans = AlltoAllTrans(classes) |
| 374 | |
| 375 | elif attack == "all-to-single": |
| 376 | indices = np.arange(len(labels)) |
| 377 | attack_trans = AlltoSingleTrans(tar_label) |
| 378 | |
| 379 | else: |
| 380 | raise ValueError('backdoor attack {} is not supported'.format(attack)) |
| 381 | |
| 382 | backdoor_num = int(len(indices) * proportion) |
| 383 | backdoor_indices = np.random.permutation(indices)[:backdoor_num] |
| 384 | |
| 385 | return backdoor_indices, trigger_trans, attack_trans |
| 386 | |
| 387 | |
| 388 | class BackdoorDataset: |
no test coverage detected