# We can use an image folder dataset the way we have it setup. # Create the dataset [dataset](https://docs.pytorch.org/vision/stable/generated/torchvision.datasets.ImageFolder.html#torchvision.datasets.ImageFolder "torchvision.datasets.ImageFolder") = [dset.ImageFolder](https://docs.pytorch.org/vision/stable/generated/torchvision.datasets.ImageFolder.html#torchvision.datasets.ImageFolder "torchvision.datasets.ImageFolder")(root=dataroot, transform=[transforms.Compose](https://docs.pytorch.org/vision/stable/generated/torchvision.transforms.Compose.html#torchvision.transforms.Compose "torchvision.transforms.Compose")([ [transforms.Resize](https://docs.pytorch.org/vision/stable/generated/torchvision.transforms.Resize.html#torchvision.transforms.Resize "torchvision.transforms.Resize")(image_size), [transforms.CenterCrop](https://docs.pytorch.org/vision/stable/generated/torchvision.transforms.CenterCrop.html#torchvision.transforms.CenterCrop "torchvision.transforms.CenterCrop")(image_size), [transforms.ToTensor](https://docs.pytorch.org/vision/stable/generated/torchvision.transforms.ToTensor.html#torchvision.transforms.ToTensor "torchvision.transforms.ToTensor")(), [transforms.Normalize](https://docs.pytorch.org/vision/stable/generated/torchvision.transforms.Normalize.html#torchvision.transforms.Normalize "torchvision.transforms.Normalize")((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)), ])) # Create the dataloader [dataloader](https://docs.pytorch.org/docs/stable/data.html#torch.utils.data.DataLoader "torch.utils.data.DataLoader") = [torch.utils.data.DataLoader](https://docs.pytorch.org/docs/stable/data.html#torch.utils.data.DataLoader "torch.utils.data.DataLoader")([dataset](https://docs.pytorch.org/vision/stable/generated/torchvision.datasets.ImageFolder.html#torchvision.datasets.ImageFolder "torchvision.datasets.ImageFolder"), batch_size=batch_size, shuffle=True, num_workers=workers) # Decide which device we want to run on [device](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.device "torch.device") = [torch.device](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.device "torch.device")("cuda:0" if ([torch.cuda.is_available](https://docs.pytorch.org/docs/stable/generated/torch.cuda.is_available.html#torch.cuda.is_available "torch.cuda.is_available")() and ngpu > 0) else "cpu") # Plot some training images real_batch = next(iter([dataloader](https://docs.pytorch.org/docs/stable/data.html#torch.utils.data.DataLoader "torch.utils.data.DataLoader"))) plt.figure(figsize=(8,8)) plt.axis("off") plt.title("Training Images") plt.imshow(np.transpose([vutils.make_grid](https://docs.pytorch.org/vision/stable/generated/torchvision.utils.make_grid.html#torchvision.utils.make_grid "torchvision.utils.make_grid")(real_batch[0].to([device](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.device "torch.device"))[:64], padding=2, normalize=True).cpu(),(1,2,0))) plt.show()