| 35 | |
| 36 | |
| 37 | class IndexedTensorDataset(Dataset): |
| 38 | def __init__(self, dataset): |
| 39 | assert isinstance(dataset, Dataset) |
| 40 | self.loader = dataset.loader |
| 41 | self.classes = dataset.classes |
| 42 | self.samples = dataset.samples |
| 43 | self.y = dataset.y |
| 44 | self.transform = transforms.Compose([ transforms.Resize( [256, 256] ) ]) |
| 45 | self.target_transform = None |
| 46 | |
| 47 | def __getitem__(self, idx): |
| 48 | x, y = super().__getitem__(idx) |
| 49 | ''' transform HWC pic to CWH pic ''' |
| 50 | x = np.asarray(x, dtype=np.uint8) |
| 51 | x = torch.tensor(x, dtype=torch.float32).permute(2,0,1) |
| 52 | return x, y, idx |
| 53 | |
| 54 | def __len__(self): |
| 55 | return len(self.y) |
| 56 | |
| 57 | |
| 58 | class PoisonedDataset(Dataset): |
no outgoing calls
no test coverage detected