| 19 | |
| 20 | |
| 21 | def get_loader(image_size): |
| 22 | transform = transforms.Compose( |
| 23 | [ |
| 24 | transforms.Resize((image_size, image_size)), |
| 25 | transforms.ToTensor(), |
| 26 | transforms.RandomHorizontalFlip(p=0.5), |
| 27 | transforms.Normalize( |
| 28 | [0.5 for _ in range(CHANNELS_IMG)], |
| 29 | [0.5 for _ in range(CHANNELS_IMG)], |
| 30 | ), |
| 31 | ] |
| 32 | ) |
| 33 | batch_size = BATCH_SIZES[int(log2(image_size / 4))] |
| 34 | dataset = datasets.ImageFolder(root=DATASET, transform=transform) |
| 35 | loader = DataLoader( |
| 36 | dataset, |
| 37 | batch_size=batch_size, |
| 38 | shuffle=True, |
| 39 | num_workers=NUM_WORKERS, |
| 40 | pin_memory=True, |
| 41 | ) |
| 42 | return loader, dataset |
| 43 | |
| 44 | |
| 45 | def train_fn( |