CoolFace
Apppublic

JimmyChin1998/Pytorch-Learning-File

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
data_setup.py66 linesDownload Raw Back to root
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