from PIL import Image from torchvision import transforms from .config import IMAGENET_MEAN, IMAGENET_STD, MODEL_CONFIG def get_transform() -> transforms.Compose: size = MODEL_CONFIG["img_size"] return transforms.Compose( [ transforms.Resize((size, size)), transforms.ToTensor(), transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD), ] ) def preprocess_image(image: Image.Image) -> Image.Image: """Ensure image is RGB and ready for the transform pipeline.""" if image.mode != "RGB": return image.convert("RGB") return image