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)), # keep 512p transforms.Grayscale(num_output_channels=3), # convert to 3-ch for pre-trained ViT transforms.ToTensor() # handle scaling # Remove Normalize since doing reconstruction, Autoencoder -> target to be in same range as output (0,1) #transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) # Std ViT normalization (ImageNet) ]) # Train only on 'normal' images 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