JimmyChin1998/Pytorch-Learning-File
0
1"""
2Contains functionality for creating PyTorch DataLoaders for
3image classification data.
4"""
5import os
6
7from torchvision import datasets, transforms
8from torch.utils.data import DataLoader
9
10NUM_WORKERS = os.cpu_count()
11
12def create_dataloaders(
13 train_dir: str,
14 test_dir: str,
15 transform: transforms.Compose,
16 batch_size: int,
17 num_workers: int=NUM_WORKERS
18):
19 """Creates training and testing DataLoaders.
20
21 Takes in a training directory and testing directory path and turns
22 them into PyTorch Datasets and then into PyTorch DataLoaders.
23
24 Args:
25 train_dir: Path to training directory.
26 test_dir: Path to testing directory.
27 transform: torchvision transforms to perform on training and testing data.
28 batch_size: Number of samples per batch in each of the DataLoaders.
29 num_workers: An integer for number of workers per DataLoader.
30
31 Returns:
32 A tuple of (train_dataloader, test_dataloader, class_names).
33 Where class_names is a list of the target classes.
34 Example usage:
35 train_dataloader, test_dataloader, class_names = \
36 = create_dataloaders(train_dir=path/to/train_dir,
37 test_dir=path/to/test_dir,
38 transform=some_transform,
39 batch_size=32,
40 num_workers=4)
41 """
42 # Use ImageFolder to create dataset(s)
43 train_data = datasets.ImageFolder(train_dir, transform=transform)
44 test_data = datasets.ImageFolder(test_dir, transform=transform)
45
46 # Get class names
47 class_names = train_data.classes
48
49 # Turn images into data loaders
50 train_dataloader = DataLoader(
51 train_data,
52 batch_size=batch_size,
53 shuffle=True,
54 num_workers=num_workers,
55 pin_memory=True,
56 )
57 test_dataloader = DataLoader(
58 test_data,
59 batch_size=batch_size,
60 shuffle=False, # don't need to shuffle test data
61 num_workers=num_workers,
62 pin_memory=True,
63 )
64
65 return train_dataloader, test_dataloader, class_names
66 