import torch, torchvision from torchvision import transforms g = torch.Generator().manual_seed(0) row_perm = torch.randperm(32, generator=g) norm = transforms.Normalize([0.5] * 3, [0.5] * 3) base = [transforms.ToTensor(), norm] shuffled = base + [transforms.Lambda(lambda x: x[:, row_perm, :])] train_nat = torchvision.datasets.CIFAR10(root="./data", train=True, download=True, transform=transforms.Compose(base)) train_shuf = torchvision.datasets.CIFAR10(root="./data", train=True, download=True, transform=transforms.Compose(shuffled)) # ...and the same two for train=False