| 63 | |
| 64 | |
| 65 | class IterableImageDataset(data.IterableDataset): |
| 66 | |
| 67 | def __init__( |
| 68 | self, |
| 69 | root, |
| 70 | parser=None, |
| 71 | split='train', |
| 72 | is_training=False, |
| 73 | batch_size=None, |
| 74 | class_map='', |
| 75 | load_bytes=False, |
| 76 | repeats=0, |
| 77 | transform=None, |
| 78 | ): |
| 79 | assert parser is not None |
| 80 | if isinstance(parser, str): |
| 81 | self.parser = create_parser( |
| 82 | parser, root=root, split=split, is_training=is_training, batch_size=batch_size, repeats=repeats) |
| 83 | else: |
| 84 | self.parser = parser |
| 85 | self.transform = transform |
| 86 | self._consecutive_errors = 0 |
| 87 | |
| 88 | def __iter__(self): |
| 89 | for img, target in self.parser: |
| 90 | if self.transform is not None: |
| 91 | img = self.transform(img) |
| 92 | if target is None: |
| 93 | target = torch.tensor(-1, dtype=torch.long) |
| 94 | yield img, target |
| 95 | |
| 96 | def __len__(self): |
| 97 | if hasattr(self.parser, '__len__'): |
| 98 | return len(self.parser) |
| 99 | else: |
| 100 | return 0 |
| 101 | |
| 102 | def filename(self, index, basename=False, absolute=False): |
| 103 | assert False, 'Filename lookup by index not supported, use filenames().' |
| 104 | |
| 105 | def filenames(self, basename=False, absolute=False): |
| 106 | return self.parser.filenames(basename, absolute) |
| 107 | |
| 108 | |
| 109 | class AugMixDataset(torch.utils.data.Dataset): |