import torch from torchvision.transforms import v2 def get_train_transforms(image_size=224): return v2.Compose([ v2.RandomResizedCrop(size=(image_size, image_size), antialias=True), v2.RandomHorizontalFlip(p=0.5), v2.RandAugment(num_ops=2, magnitude=9), v2.ToImage(), v2.ToDtype(torch.float32, scale=True), v2.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) def get_eval_transforms(image_size=224): return v2.Compose([ v2.Resize(size=(256, 256), antialias=True), v2.CenterCrop(size=(image_size, image_size)), v2.ToImage(), v2.ToDtype(torch.float32, scale=True), v2.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) def get_batch_synthetics(num_classes): return v2.RandomChoice([ v2.CutMix(num_classes=num_classes, alpha=1.0), v2.MixUp(num_classes=num_classes, alpha=0.8) ])