folder.py 10.8 KB
Newer Older
1
from .vision import VisionDataset
soumith's avatar
soumith committed
2
3

from PIL import Image
4

soumith's avatar
soumith committed
5
6
import os
import os.path
Philip Meier's avatar
Philip Meier committed
7
from typing import Any, Callable, cast, Dict, List, Optional, Tuple
soumith's avatar
soumith committed
8

9

Philip Meier's avatar
Philip Meier committed
10
def has_file_allowed_extension(filename: str, extensions: Tuple[str, ...]) -> bool:
11
    """Checks if a file is an allowed extension.
12
13
14

    Args:
        filename (string): path to a file
15
        extensions (tuple of strings): extensions to consider (lowercase)
16
17

    Returns:
18
        bool: True if the filename ends with one of given extensions
19
    """
20
    return filename.lower().endswith(extensions)
soumith's avatar
soumith committed
21

22

Philip Meier's avatar
Philip Meier committed
23
def is_image_file(filename: str) -> bool:
24
25
26
27
28
29
30
31
32
33
34
    """Checks if a file is an allowed image extension.

    Args:
        filename (string): path to a file

    Returns:
        bool: True if the filename ends with a known image extension
    """
    return has_file_allowed_extension(filename, IMG_EXTENSIONS)


35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
def find_classes(directory: str) -> Tuple[List[str], Dict[str, int]]:
    """Finds the class folders in a dataset structured as follows:

    .. code::

        directory/
        ├── class_x
        │   ├── xxx.ext
        │   ├── xxy.ext
        │   └── ...
        │       └── xxz.ext
        └── class_y
            ├── 123.ext
            ├── nsdf3.ext
            └── ...
                └── asd932_.ext

    Args:
        directory (str): Root directory path.

    Raises:
        FileNotFoundError: If ``directory`` has no class folders.

    Returns:
        (Tuple[List[str], Dict[str, int]]): List of all classes and dictionary mapping each class to an index.
    """
    classes = sorted(entry.name for entry in os.scandir(directory) if entry.is_dir())
    if not classes:
        raise FileNotFoundError(f"Couldn't find any class folder in {directory}.")

    class_to_idx = {cls_name: i for i, cls_name in enumerate(classes)}
    return classes, class_to_idx


Philip Meier's avatar
Philip Meier committed
69
70
def make_dataset(
    directory: str,
71
    class_to_idx: Optional[Dict[str, int]] = None,
Philip Meier's avatar
Philip Meier committed
72
73
74
    extensions: Optional[Tuple[str, ...]] = None,
    is_valid_file: Optional[Callable[[str], bool]] = None,
) -> List[Tuple[str, int]]:
75
76
77
78
    """Generates a list of samples of a form (path_to_sample, class).

    Args:
        directory (str): root dataset directory
79
80
        class_to_idx (Optional[Dict[str, int]]): Dictionary mapping class name to class index. If omitted, is generated
            by :func:`find_classes`.
81
82
83
84
85
86
87
88
        extensions (optional): A list of allowed extensions.
            Either extensions or is_valid_file should be passed. Defaults to None.
        is_valid_file (optional): A function that takes path of a file
            and checks if the file is a valid file
            (used to check of corrupt files) both extensions and
            is_valid_file should not be passed. Defaults to None.

    Raises:
89
        ValueError: In case ``class_to_idx`` is empty.
90
        ValueError: In case ``extensions`` and ``is_valid_file`` are None or both are not None.
91
        FileNotFoundError: In case no valid file was found for any class.
92
93
94
95

    Returns:
        List[Tuple[str, int]]: samples of a form (path_to_sample, class)
    """
96
    directory = os.path.expanduser(directory)
97
98
99
100
101
102

    if class_to_idx is None:
        _, class_to_idx = find_classes(directory)
    elif not class_to_idx:
        raise ValueError("'class_to_index' must have at least one entry to collect any samples.")

103
104
105
    both_none = extensions is None and is_valid_file is None
    both_something = extensions is not None and is_valid_file is not None
    if both_none or both_something:
Surgan Jandial's avatar
Surgan Jandial committed
106
        raise ValueError("Both extensions and is_valid_file cannot be None or not None at the same time")
107

108
    if extensions is not None:
109

Philip Meier's avatar
Philip Meier committed
110
111
        def is_valid_file(x: str) -> bool:
            return has_file_allowed_extension(x, cast(Tuple[str, ...], extensions))
112

Philip Meier's avatar
Philip Meier committed
113
    is_valid_file = cast(Callable[[str], bool], is_valid_file)
114
115
116

    instances = []
    available_classes = set()
117
118
119
120
    for target_class in sorted(class_to_idx.keys()):
        class_index = class_to_idx[target_class]
        target_dir = os.path.join(directory, target_class)
        if not os.path.isdir(target_dir):
soumith's avatar
soumith committed
121
            continue
122
        for root, _, fnames in sorted(os.walk(target_dir, followlinks=True)):
123
            for fname in sorted(fnames):
124
125
                path = os.path.join(root, fname)
                if is_valid_file(path):
126
127
                    item = path, class_index
                    instances.append(item)
128
129
130
131

                    if target_class not in available_classes:
                        available_classes.add(target_class)

132
    empty_classes = set(class_to_idx.keys()) - available_classes
133
134
135
136
137
138
    if empty_classes:
        msg = f"Found no valid file for the classes {', '.join(sorted(empty_classes))}. "
        if extensions is not None:
            msg += f"Supported extensions are: {', '.join(extensions)}"
        raise FileNotFoundError(msg)

139
    return instances
soumith's avatar
soumith committed
140

141

142
class DatasetFolder(VisionDataset):
143
144
145
146
    """A generic data loader where the samples are arranged in this way: ::

        root/class_x/xxx.ext
        root/class_x/xxy.ext
147
        root/class_x/[...]/xxz.ext
148
149
150

        root/class_y/123.ext
        root/class_y/nsdf3.ext
151
        root/class_y/[...]/asd932_.ext
152
153
154
155

    Args:
        root (string): Root directory path.
        loader (callable): A function to load a sample given its path.
156
        extensions (tuple[string]): A list of allowed extensions.
157
            both extensions and is_valid_file should not be passed.
158
159
160
161
162
        transform (callable, optional): A function/transform that takes in
            a sample and returns a transformed version.
            E.g, ``transforms.RandomCrop`` for images.
        target_transform (callable, optional): A function/transform that takes
            in the target and transforms it.
Carrie's avatar
Carrie committed
163
164
        is_valid_file (callable, optional): A function that takes path of a file
            and check if the file is a valid file (used to check of corrupt files)
165
            both extensions and is_valid_file should not be passed.
166
167

     Attributes:
168
        classes (list): List of the class names sorted alphabetically.
169
170
        class_to_idx (dict): Dict with items (class_name, class_index).
        samples (list): List of (sample path, class_index) tuples
171
        targets (list): The class_index value for each image in the dataset
172
173
    """

Philip Meier's avatar
Philip Meier committed
174
175
176
177
178
179
180
181
182
    def __init__(
            self,
            root: str,
            loader: Callable[[str], Any],
            extensions: Optional[Tuple[str, ...]] = None,
            transform: Optional[Callable] = None,
            target_transform: Optional[Callable] = None,
            is_valid_file: Optional[Callable[[str], bool]] = None,
    ) -> None:
183
184
        super(DatasetFolder, self).__init__(root, transform=transform,
                                            target_transform=target_transform)
185
        classes, class_to_idx = self.find_classes(self.root)
186
        samples = self.make_dataset(self.root, class_to_idx, extensions, is_valid_file)
187
188
189
190
191
192
193

        self.loader = loader
        self.extensions = extensions

        self.classes = classes
        self.class_to_idx = class_to_idx
        self.samples = samples
194
        self.targets = [s[1] for s in samples]
195

196
197
198
199
200
201
202
203
204
    @staticmethod
    def make_dataset(
        directory: str,
        class_to_idx: Dict[str, int],
        extensions: Optional[Tuple[str, ...]] = None,
        is_valid_file: Optional[Callable[[str], bool]] = None,
    ) -> List[Tuple[str, int]]:
        return make_dataset(directory, class_to_idx, extensions=extensions, is_valid_file=is_valid_file)

205
206
207
208
209
210
    def find_classes(self, dir: str) -> Tuple[List[str], Dict[str, int]]:
        """Same as :func:`find_classes`.

        This method can be overridden to only consider
        a subset of classes, or to adapt to a different dataset directory structure.
        """
211
        return find_classes(dir)
212

Philip Meier's avatar
Philip Meier committed
213
    def __getitem__(self, index: int) -> Tuple[Any, Any]:
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
        """
        Args:
            index (int): Index

        Returns:
            tuple: (sample, target) where target is class_index of the target class.
        """
        path, target = self.samples[index]
        sample = self.loader(path)
        if self.transform is not None:
            sample = self.transform(sample)
        if self.target_transform is not None:
            target = self.target_transform(target)

        return sample, target

Philip Meier's avatar
Philip Meier committed
230
    def __len__(self) -> int:
231
232
233
        return len(self.samples)


234
IMG_EXTENSIONS = ('.jpg', '.jpeg', '.png', '.ppm', '.bmp', '.pgm', '.tif', '.tiff', '.webp')
235
236


Philip Meier's avatar
Philip Meier committed
237
def pil_loader(path: str) -> Image.Image:
238
239
    # open path as file to avoid ResourceWarning (https://github.com/python-pillow/Pillow/issues/835)
    with open(path, 'rb') as f:
240
241
        img = Image.open(f)
        return img.convert('RGB')
242
243


Philip Meier's avatar
Philip Meier committed
244
245
# TODO: specify the return type
def accimage_loader(path: str) -> Any:
246
247
248
249
250
251
252
253
    import accimage
    try:
        return accimage.Image(path)
    except IOError:
        # Potentially a decoding problem, fall back to PIL.Image
        return pil_loader(path)


Philip Meier's avatar
Philip Meier committed
254
def default_loader(path: str) -> Any:
255
256
257
258
259
260
261
    from torchvision import get_image_backend
    if get_image_backend() == 'accimage':
        return accimage_loader(path)
    else:
        return pil_loader(path)


262
class ImageFolder(DatasetFolder):
263
264
265
266
    """A generic data loader where the images are arranged in this way: ::

        root/dog/xxx.png
        root/dog/xxy.png
267
        root/dog/[...]/xxz.png
268
269
270

        root/cat/123.png
        root/cat/nsdf3.png
271
        root/cat/[...]/asd932_.png
272
273
274
275
276
277
278
279

    Args:
        root (string): Root directory path.
        transform (callable, optional): A function/transform that  takes in an PIL image
            and returns a transformed version. E.g, ``transforms.RandomCrop``
        target_transform (callable, optional): A function/transform that takes in the
            target and transforms it.
        loader (callable, optional): A function to load an image given its path.
280
        is_valid_file (callable, optional): A function that takes path of an Image file
Carrie's avatar
Carrie committed
281
            and check if the file is a valid file (used to check of corrupt files)
282
283

     Attributes:
284
        classes (list): List of the class names sorted alphabetically.
285
286
287
        class_to_idx (dict): Dict with items (class_name, class_index).
        imgs (list): List of (image path, class_index) tuples
    """
288

Philip Meier's avatar
Philip Meier committed
289
290
291
292
293
294
295
296
    def __init__(
            self,
            root: str,
            transform: Optional[Callable] = None,
            target_transform: Optional[Callable] = None,
            loader: Callable[[str], Any] = default_loader,
            is_valid_file: Optional[Callable[[str], bool]] = None,
    ):
297
        super(ImageFolder, self).__init__(root, loader, IMG_EXTENSIONS if is_valid_file is None else None,
298
                                          transform=transform,
299
300
                                          target_transform=target_transform,
                                          is_valid_file=is_valid_file)
301
        self.imgs = self.samples