| 56 | |
| 57 | |
| 58 | class PoisonedDataset(Dataset): |
| 59 | def __init__(self, dataset, noise, fitr=None): |
| 60 | assert isinstance(dataset, Dataset) |
| 61 | self.loader = dataset.loader |
| 62 | self.classes = dataset.classes |
| 63 | self.samples = dataset.samples |
| 64 | self.y = dataset.y |
| 65 | self.transform = transforms.Compose([ transforms.Resize([ noise.shape[1], noise.shape[2] ]) ]) |
| 66 | self.target_transform = None |
| 67 | self.data_transform = dataset.transform |
| 68 | self.data_fitr = fitr |
| 69 | ''' the shape of the noise should be (NHWC) ''' |
| 70 | self.noise = noise |
| 71 | |
| 72 | def __getitem__(self, idx): |
| 73 | x, y = super().__getitem__(idx) |
| 74 | x = (np.asarray(x, dtype=np.int16) + self.noise[idx].astype(np.int16)).clip(0, 255).astype(np.uint8) |
| 75 | |
| 76 | ''' low pass filtering ''' |
| 77 | if self.data_fitr is not None: |
| 78 | x = self.data_fitr(x) |
| 79 | |
| 80 | x = self.data_transform( Image.fromarray(x) ) |
| 81 | # x = self.data_transform( Image.fromarray(x.astype(np.uint8)) ) |
| 82 | return x, y |
| 83 | |
| 84 | def __len__(self): |
| 85 | return len(self.y) |
| 86 | |
| 87 | |
| 88 | class IndexedDataset(Dataset): |
no outgoing calls
no test coverage detected