File size: 2,537 Bytes
228add1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
"""
Inference image preprocessing.
Handles loading images from file path, URL, or base64 string.
"""

import base64
from io import BytesIO
from pathlib import Path
from typing import Union

import requests
from PIL import Image
from torchvision import transforms

from ml.src.data.transforms import IMAGENET_MEAN, IMAGENET_STD


def load_image_from_path(path: str | Path) -> Image.Image:
    """Load an image from a file path."""
    return Image.open(path).convert('RGB')


def load_image_from_url(url: str, timeout: int = 10) -> Image.Image:
    """Load an image from a URL."""
    response = requests.get(url, timeout=timeout, stream=True)
    response.raise_for_status()
    return Image.open(BytesIO(response.content)).convert('RGB')


def load_image_from_base64(b64_string: str) -> Image.Image:
    """Load an image from a base64-encoded string."""
    # Remove data URI prefix if present
    if ',' in b64_string:
        b64_string = b64_string.split(',', 1)[1]
    image_bytes = base64.b64decode(b64_string)
    return Image.open(BytesIO(image_bytes)).convert('RGB')


def load_image_from_bytes(raw_bytes: bytes) -> Image.Image:
    """Load an image from raw bytes."""
    return Image.open(BytesIO(raw_bytes)).convert('RGB')


def get_inference_transform(img_size: int = 224) -> transforms.Compose:
    """Standard inference preprocessing pipeline."""
    return transforms.Compose([
        transforms.Resize((img_size, img_size)),
        transforms.ToTensor(),
        transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),
    ])


def preprocess_image(
    image: Union[str, bytes, Image.Image],
    img_size: int = 224,
) -> tuple:
    """
    Universal image preprocessor.
    Accepts path string, URL, base64, bytes, or PIL Image.
    Returns (tensor, pil_image).
    """
    # Load image if needed
    if isinstance(image, str):
        if image.startswith(('http://', 'https://')):
            pil_image = load_image_from_url(image)
        elif image.startswith('data:image') or len(image) > 500:
            pil_image = load_image_from_base64(image)
        else:
            pil_image = load_image_from_path(image)
    elif isinstance(image, bytes):
        pil_image = load_image_from_bytes(image)
    elif isinstance(image, Image.Image):
        pil_image = image.convert('RGB')
    else:
        raise ValueError(f"Unsupported image type: {type(image)}")

    transform = get_inference_transform(img_size)
    tensor = transform(pil_image).unsqueeze(0)  # Add batch dimension

    return tensor, pil_image