| 57 | |
| 58 | |
| 59 | class ImageFolder(data.Dataset): |
| 60 | |
| 61 | def __init__(self, root, transform=None, return_paths=False, |
| 62 | loader=default_loader): |
| 63 | imgs = make_dataset(root) |
| 64 | if len(imgs) == 0: |
| 65 | raise(RuntimeError("Found 0 images in: " + root + "\n" |
| 66 | "Supported image extensions are: " + |
| 67 | ",".join(IMG_EXTENSIONS))) |
| 68 | |
| 69 | self.root = root |
| 70 | self.imgs = imgs |
| 71 | self.transform = transform |
| 72 | self.return_paths = return_paths |
| 73 | self.loader = loader |
| 74 | |
| 75 | def __getitem__(self, index): |
| 76 | path = self.imgs[index] |
| 77 | img = self.loader(path) |
| 78 | if self.transform is not None: |
| 79 | img = self.transform(img) |
| 80 | if self.return_paths: |
| 81 | return img, path |
| 82 | else: |
| 83 | return img |
| 84 | |
| 85 | def __len__(self): |
| 86 | return len(self.imgs) |
nothing calls this directly
no outgoing calls
no test coverage detected