from torchvision.transforms import v2 as T def get_transform(train): transforms = [] if train: transforms.append([T.RandomHorizontalFlip](https://docs.pytorch.org/vision/stable/generated/torchvision.transforms.v2.RandomHorizontalFlip.html#torchvision.transforms.v2.RandomHorizontalFlip "torchvision.transforms.v2.RandomHorizontalFlip")(0.5)) transforms.append([T.ToDtype](https://docs.pytorch.org/vision/stable/generated/torchvision.transforms.v2.ToDtype.html#torchvision.transforms.v2.ToDtype "torchvision.transforms.v2.ToDtype")([torch.float](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.dtype "torch.dtype"), scale=True)) transforms.append([T.ToPureTensor](https://docs.pytorch.org/vision/stable/generated/torchvision.transforms.v2.ToPureTensor.html#torchvision.transforms.v2.ToPureTensor "torchvision.transforms.v2.ToPureTensor")()) return [T.Compose](https://docs.pytorch.org/vision/stable/generated/torchvision.transforms.v2.Compose.html#torchvision.transforms.v2.Compose "torchvision.transforms.v2.Compose")(transforms)