| import torch |
| from torchvision import datasets, transforms |
| from torch.utils.data import DataLoader |
|
|
| def create_dataloaders(data_dir, batch_size=8, num_workers=2): |
| transform = transforms.Compose([ |
| transforms.Resize((512, 512)), |
| transforms.Grayscale(num_output_channels=3), |
| transforms.ToTensor() |
| |
| |
| ]) |
|
|
| |
| train_dataset = datasets.ImageFolder(root=f"{data_dir}/train", transform=transform) |
| test_dataset = datasets.ImageFolder(root=f"{data_dir}/test", transform=transform) |
|
|
| train_loader = DataLoader(train_dataset, batch_size=batch_size, num_workers=num_workers, shuffle=True) |
| test_loader = DataLoader(test_dataset, batch_size=batch_size, num_workers=num_workers, shuffle=False) |
|
|
| return train_loader, test_loader |