File size: 3,748 Bytes
eeeea78
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
import numpy as np
from PIL import Image
import io
from typing import Tuple, Dict, Union
import asyncio
from concurrent.futures import ThreadPoolExecutor

# Optimized thread pool for CPU-intensive image operations
image_executor = ThreadPoolExecutor(max_workers=2, thread_name_prefix="ImageProcessor")

# Common model input sizes
MODEL_SIZES = {
    'xceptionnet': (299, 299),
    'mesonet': (256, 256)
}


def _preprocess_image_core(image_data: bytes, target_size: Tuple[int, int]) -> np.ndarray:
    """
    Core image preprocessing function (synchronous).
    
    Args:
        image_data: Raw image bytes
        target_size: Target size as (width, height)
    
    Returns:
        Preprocessed numpy array ready for model prediction
    """
    # Read and validate image
    image = Image.open(io.BytesIO(image_data))
    
    # Convert to RGB if needed
    if image.mode != 'RGB':
        image = image.convert('RGB')
    
    # Resize and normalize
    image = image.resize(target_size, Image.Resampling.LANCZOS)
    img_array = np.array(image, dtype=np.float32) / 255.0
    
    # Add batch dimension
    return np.expand_dims(img_array, axis=0)


async def preprocess_image(image_data: bytes, target_size: Tuple[int, int]) -> np.ndarray:
    """
    Async image preprocessing for FastAPI.
    
    Args:
        image_data: Raw image bytes from uploaded file
        target_size: Target size as (width, height)
    
    Returns:
        Preprocessed numpy array ready for model prediction
    """
    loop = asyncio.get_event_loop()
    return await loop.run_in_executor(image_executor, _preprocess_image_core, image_data, target_size)


async def preprocess_for_models(image_data: bytes) -> Dict[str, np.ndarray]:
    """
    Preprocess image for all models simultaneously (optimized).
    
    Args:
        image_data: Raw image bytes from uploaded file
    
    Returns:
        Dictionary with preprocessed arrays for each model
    """
    def _process_all_sizes(data: bytes) -> Dict[str, np.ndarray]:
        # Load image once
        image = Image.open(io.BytesIO(data))
        if image.mode != 'RGB':
            image = image.convert('RGB')
        
        results = {}
        for model_name, size in MODEL_SIZES.items():
            # Resize and normalize
            resized = image.resize(size, Image.Resampling.LANCZOS)
            img_array = np.array(resized, dtype=np.float32) / 255.0
            results[model_name] = np.expand_dims(img_array, axis=0)
        
        return results
    
    loop = asyncio.get_event_loop()
    return await loop.run_in_executor(image_executor, _process_all_sizes, image_data)


def validate_image(image_data: bytes) -> bool:
    """
    Validate if the provided bytes represent a valid image.
    
    Args:
        image_data: Raw image bytes
    
    Returns:
        True if valid image, False otherwise
    """
    try:
        with Image.open(io.BytesIO(image_data)) as img:
            img.verify()
        return True
    except Exception:
        return False


# Legacy compatibility function (for any remaining Flask code)
def preprocess_image_sync(file_or_bytes: Union[bytes, object], target_size: Tuple[int, int]) -> np.ndarray:
    """Legacy synchronous preprocessing function for backward compatibility."""
    if isinstance(file_or_bytes, bytes):
        return _preprocess_image_core(file_or_bytes, target_size)
    else:
        # Assume it's a file-like object
        image = Image.open(file_or_bytes)
        if image.mode != 'RGB':
            image = image.convert('RGB')
        
        image = image.resize(target_size, Image.Resampling.LANCZOS)
        img_array = np.array(image, dtype=np.float32) / 255.0
        return np.expand_dims(img_array, axis=0)