| 28 | |
| 29 | |
| 30 | def get_loader(image_size): |
| 31 | transform = transforms.Compose( |
| 32 | [ |
| 33 | transforms.Resize((image_size, image_size)), |
| 34 | transforms.ToTensor(), |
| 35 | transforms.RandomHorizontalFlip(p=0.5), |
| 36 | transforms.Normalize( |
| 37 | [0.5 for _ in range(config.CHANNELS_IMG)], |
| 38 | [0.5 for _ in range(config.CHANNELS_IMG)], |
| 39 | ), |
| 40 | ] |
| 41 | ) |
| 42 | batch_size = config.BATCH_SIZES[int(log2(image_size / 4))] |
| 43 | dataset = datasets.ImageFolder(root=config.DATASET, transform=transform) |
| 44 | loader = DataLoader( |
| 45 | dataset, |
| 46 | batch_size=batch_size, |
| 47 | shuffle=True, |
| 48 | num_workers=config.NUM_WORKERS, |
| 49 | pin_memory=True, |
| 50 | ) |
| 51 | return loader, dataset |
| 52 | |
| 53 | |
| 54 | def train_fn( |