| 44 | |
| 45 | |
| 46 | class FakeLabelDataset(Dataset): |
| 47 | def __init__(self, dataset, root=None, transform=None): |
| 48 | super(FakeLabelDataset, self).__init__() |
| 49 | self.dataset = dataset |
| 50 | self.root = root |
| 51 | self.transform = transform |
| 52 | if isinstance(self.dataset[0][0], str): |
| 53 | self.data0_is_numpy = False |
| 54 | else: |
| 55 | self.data0_is_numpy = True |
| 56 | def __len__(self): |
| 57 | return len(self.dataset) |
| 58 | |
| 59 | def __getitem__(self, indices): |
| 60 | return self._get_single_item(indices) |
| 61 | |
| 62 | def _get_single_item(self, index): |
| 63 | if self.data0_is_numpy: |
| 64 | image_numpy, fake, real = self.dataset[index] |
| 65 | |
| 66 | img = Image.fromarray(image_numpy.astype('uint8')).convert('RGB') |
| 67 | if self.transform is not None: |
| 68 | img = self.transform(img) |
| 69 | |
| 70 | return img, fake, real |
| 71 | |
| 72 | else: |
| 73 | fname, fake, real = self.dataset[index] |
| 74 | fpath = fname |
| 75 | if self.root is not None: |
| 76 | fpath = osp.join(self.root, fname) |
| 77 | |
| 78 | img = Image.open(fpath).convert('RGB') |
| 79 | |
| 80 | if self.transform is not None: |
| 81 | img = self.transform(img) |
| 82 | |
| 83 | return img, fake, real |