(self,
path, # Path to directory or zip.
resolution = None, # Ensure specific resolution, None = highest available.
**super_kwargs, # Additional arguments for the Dataset base class.
)
| 155 | |
| 156 | class ImageFolderDataset(Dataset): |
| 157 | def __init__(self, |
| 158 | path, # Path to directory or zip. |
| 159 | resolution = None, # Ensure specific resolution, None = highest available. |
| 160 | **super_kwargs, # Additional arguments for the Dataset base class. |
| 161 | ): |
| 162 | self._path = path |
| 163 | self._zipfile = None |
| 164 | |
| 165 | if os.path.isdir(self._path): |
| 166 | self._type = 'dir' |
| 167 | self._all_fnames = {os.path.relpath(os.path.join(root, fname), start=self._path) for root, _dirs, files in os.walk(self._path) for fname in files} |
| 168 | elif self._file_ext(self._path) == '.zip': |
| 169 | self._type = 'zip' |
| 170 | self._all_fnames = set(self._get_zipfile().namelist()) |
| 171 | else: |
| 172 | raise IOError('Path must point to a directory or zip') |
| 173 | |
| 174 | PIL.Image.init() |
| 175 | self._image_fnames = sorted(fname for fname in self._all_fnames if self._file_ext(fname) in PIL.Image.EXTENSION) |
| 176 | if len(self._image_fnames) == 0: |
| 177 | raise IOError('No image files found in the specified path') |
| 178 | |
| 179 | name = os.path.splitext(os.path.basename(self._path))[0] |
| 180 | raw_shape = [len(self._image_fnames)] + list(self._load_raw_image(0).shape) |
| 181 | if resolution is not None and (raw_shape[2] != resolution or raw_shape[3] != resolution): |
| 182 | raise IOError('Image files do not match the specified resolution') |
| 183 | super().__init__(name=name, raw_shape=raw_shape, **super_kwargs) |
| 184 | |
| 185 | @staticmethod |
| 186 | def _file_ext(fname): |
nothing calls this directly
no test coverage detected