from monai.transforms import ( Compose, LoadImaged, EnsureChannelFirstd, ScaleIntensityd, AsDiscreted, Resized, RandFlipd, EnsureTyped, ) import torch base = [ LoadImaged(keys=["image", "label"], reader="PILReader", image_only=True), EnsureChannelFirstd(keys=["image", "label"]), ScaleIntensityd(keys="image"), # brightness spread AsDiscreted(keys="label", threshold=0.5), # clean {0, 1} mask Resized(keys=["image", "label"], # hundreds of sizes spatial_size=cfg.image_size, mode=("bilinear", "nearest")), ] train_transforms = Compose(base + [ RandFlipd(keys=["image", "label"], prob=0.5, spatial_axis=1), # horizontal EnsureTyped(keys=["image", "label"], dtype=torch.float32), ]) val_transforms = Compose(base + [ EnsureTyped(keys=["image", "label"], dtype=torch.float32), ])