Coverage for wrapper/yolo/wrapper_yolo.py: 69%
175 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
1import math
2import cv2
3import logging
4import os
5import numpy as np
7from typing import override
8from typing import Optional
9from typing import List
10from typing import Tuple
11from typing import Dict
12from typing import Any
13from typing import Union
15import torch
16from torch import inference_mode
18from PIL import Image
20from wrapper_model import ModelGeneration
21from image_utils import base64_to_img
23from ultralytics import YOLO
24"""
25WARNING Ultralytics settings reset to default values. This may be due to a possible problem with your settings or a
26recent ultralytics package update.
27View Ultralytics Settings with 'yolo settings' or at '/home/azureuser/.config/Ultralytics/settings.json'
28Update Settings with 'yolo settings key=value', i.e. 'yolo settings runs_dir=path/to/dir'.
29For help see https://docs.ultralytics.com/quickstart/#ultralytics-settings.
30"""
33def zoom_image_old(
34 image: Image.Image,
35 x: int,
36 y: int,
37 w: int,
38 h: int,
39 zoom_factor: float = 2
40) -> Image.Image:
41 image_np = np.array(image)
42 image_cv2 = cv2.cvtColor(image_np, cv2.COLOR_RGB2BGR)
43 w_zoom = int(image.width // zoom_factor)
44 h_zoom = int(image.height // zoom_factor)
45 x_center = x + w // 2
46 y_center = y + h // 2
47 x_zoom = int(x_center - w_zoom // 2)
48 y_zoom = int(y_center - h_zoom // 2)
50 x_zoom = max(0, x_zoom)
51 y_zoom = max(0, y_zoom)
52 x_zoom = min(image.width - w_zoom, x_zoom)
53 y_zoom = min(image.height - h_zoom, y_zoom)
54 cropped_person_zoom = image_cv2[
55 y_zoom:y_zoom + h_zoom,
56 x_zoom:x_zoom + w_zoom
57 ]
58 cropped_person_zoom_rgb = cv2.cvtColor(cropped_person_zoom, cv2.COLOR_BGR2RGB)
59 zoomed_image = Image.fromarray(cropped_person_zoom_rgb)
60 return zoomed_image
63def zoom_image(
64 image: Image.Image,
65 x: int,
66 y: int,
67 w: int,
68 h: int,
69 zoom_factor: float = 2.0
70) -> Image.Image:
71 orig_w, orig_h = image.size
72 aspect = orig_w / orig_h
74 # Initial crop size based on zoom factor
75 target_crop_w = max(w, int(orig_w / zoom_factor))
76 target_crop_h = max(h, int(orig_h / zoom_factor))
78 # Adjust crop size to match original aspect ratio
79 if target_crop_w / target_crop_h > aspect:
80 target_crop_h = int(target_crop_w / aspect)
81 else:
82 target_crop_w = int(target_crop_h * aspect)
84 # Center the crop on the box
85 center_x = x + w // 2
86 center_y = y + h // 2
88 # Compute crop box
89 left = max(0, center_x - target_crop_w // 2)
90 upper = max(0, center_y - target_crop_h // 2)
91 right = left + target_crop_w
92 lower = upper + target_crop_h
94 # Adjust if crop goes out of bounds
95 if right > orig_w:
96 overflow = right - orig_w
97 left = max(0, left - overflow)
98 right = orig_w
99 if lower > orig_h:
100 overflow = lower - orig_h
101 upper = max(0, upper - overflow)
102 lower = orig_h
104 # Final sanity check on box
105 left = int(left)
106 upper = int(upper)
107 right = int(right)
108 lower = int(lower)
110 # Crop and resize
111 image = image.crop((left, upper, right, lower))
112 # image = image.resize((orig_w, orig_h), Image.Resampling.LANCZOS)
113 return image
116def take_top_characters(
117 images: List[Tuple[int, float, Image.Image]],
118 num_characters: int
119) -> List[Image.Image]:
120 # Based on confidence
121 sorted_images = sorted(
122 images,
123 key=lambda x: x[1],
124 reverse=True
125 )[:num_characters]
127 # From left to right
128 sorted_images = sorted(
129 sorted_images,
130 key=lambda x: x[0],
131 reverse=False
132 )
133 result: List[Image.Image] = [image[2] for image in sorted_images]
134 return result
137class ImageCharacterExtractor(ModelGeneration):
138 """
139 A class to zoom into a specific area of an image.
140 """
142 def __init__(self) -> None:
143 super().__init__("yolo")
145 # Model components
146 self.obj_recognition: Optional[YOLO] = None
148 def __del__(self) -> None:
149 if self.obj_recognition is not None:
150 self.obj_recognition = None
152 def init_parallelism(self) -> None:
153 self.load_timer.start("torch_dist")
154 # No real parallelism as it runs with a single GPU or CPU
155 self.gpu: Optional[str] = None
156 device_id: Union[int, str]
157 if torch.cuda.is_available():
158 self.rank = int(os.getenv("RANK", 0))
159 self.local_rank = int(os.getenv("LOCAL_RANK", 0))
160 self.world_size = int(os.getenv("WORLD_SIZE", 1))
161 device_id = self.local_rank
162 self.device = torch.device(f"cuda:{device_id}")
163 self.gpu = torch.cuda.get_device_name(device_id)
164 torch.cuda.set_device(self.local_rank)
165 else:
166 device_id = "cpu"
167 self.device = torch.device(device_id)
168 self.device_id = device_id
169 self.load_timer.end("torch_dist")
171 def load_model(self) -> None:
172 self.load_timer.start("yolo")
173 # pretrained YOLO11n model
174 self.obj_recognition = YOLO("yolo11n.pt")
175 # TODO expand this list if we go fancier
176 self.CHARACTER_CLASSES = [
177 "person",
178 "teddy bear",
179 ]
180 self.load_timer.end("yolo")
182 def init_model_parallelism(self) -> None:
183 if self.world_size > 1:
184 logging.warning("YOLO does not support distributed parallelism.")
186 def model_compile(self) -> None:
187 if not self.torch_compile:
188 return
189 self.load_timer.start("compile")
190 assert self.obj_recognition is not None
191 self.obj_recognition.model = torch.compile(
192 self.obj_recognition.model,
193 mode="reduce-overhead")
194 self.load_timer.end("compile")
196 @inference_mode()
197 async def warmup(self) -> None:
198 logging.info("Warmup for YOLO generation")
199 empty_img = Image.new("RGB", (640, 480), (255, 255, 255))
200 self.extract_characters(empty_img)
202 def _assert_model_init(self) -> None:
203 super()._assert_model_init()
204 if self.obj_recognition is None:
205 raise ValueError("YOLO not loaded.")
207 @inference_mode()
208 def extract_characters(
209 self,
210 img: Image.Image,
211 num_characters: int = 2,
212 zoom_factor: float = 1.6,
213 job_id: Optional[str] = None,
214 ) -> List[Optional[Image.Image]]:
215 gen_timer = self._new_gen_timer(job_id)
217 self._assert_model_init()
218 assert self.obj_recognition is not None
220 self.running = True # We can run in parallel but good to know if we are running
222 try:
223 results = self.obj_recognition.predict(
224 img,
225 verbose=False,
226 device=self.device)
228 if len(results) == 0:
229 logging.warning("No results for the image.")
230 return []
232 # Original image with the boxes on top for debugging
233 debug_img_np = results[0].plot()
234 debug_img_np = cv2.cvtColor(debug_img_np, cv2.COLOR_BGR2RGB)
235 debug_img = Image.fromarray(debug_img_np)
237 if len(results[0].boxes) == 0:
238 logging.warning("No characters detected in the image.")
239 return [debug_img] + [None] * num_characters
241 logging.info(f"Detected {len(results[0].boxes)} objects in the image.")
242 person_zoom_images = []
243 for box in results[0].boxes:
244 obj_class = box.cls # class index
245 obj_class = self.obj_recognition.model.names[int(obj_class)] # It takes < 3 milliseconds
246 obj_conf = box.conf.cpu().numpy()[0] # confidence score
247 x1, y1, x2, y2 = box.xyxy.cpu().numpy().tolist()[0] # xyxy format (x1, y1, x2, y2)
248 if obj_class in self.CHARACTER_CLASSES:
249 x_center = math.floor((x2 + x1) / 2.0)
250 person_zoom_image = zoom_image(
251 img,
252 x1, y1,
253 x2 - x1, y2 - y1,
254 zoom_factor=zoom_factor)
255 person_zoom_images.append((x_center, obj_conf, person_zoom_image))
257 # Sort by confidence and take the top NUM_CHARACTERS and sort by x
258 top_images: List[Image.Image] = take_top_characters(person_zoom_images, num_characters)
259 output: List[Optional[Image.Image]] = [debug_img]
260 output.extend(top_images)
261 return output
262 finally:
263 self.running = False
264 gen_timer.end("total")
266 @override
267 @inference_mode()
268 async def generate(
269 self,
270 img: Image.Image,
271 num_characters: int = 2,
272 zoom_factor: float = 1.6,
273 job_id: Optional[str] = None,
274 ) -> List[Optional[Image.Image]]:
275 return self.extract_characters(
276 img=img,
277 num_characters=num_characters,
278 zoom_factor=zoom_factor,
279 job_id=job_id)
281 def get_health(self) -> Dict[str, Any]:
282 ret = super().get_health()
283 ret["gpu"] = self.gpu
284 return ret
286 async def get_rest_args(
287 self,
288 data_json: Dict[str, Union[str, int, float]]
289 ) -> Dict[str, Any]:
290 if data_json is None or not isinstance(data_json, dict):
291 raise ValueError("Missing JSON body")
293 job_id = data_json.get("job_id", None)
295 img_base64 = data_json.get("img", None)
296 if img_base64 is None:
297 raise ValueError("Missing 'img' parameter")
298 if not isinstance(img_base64, str):
299 raise ValueError("'img' parameter must be a base64 string")
300 img = base64_to_img(img_base64)
301 num_characters = int(data_json.get("num_characters", 2))
302 zoom_factor = float(data_json.get("zoom_factor", 1.6))
304 return {
305 "task": self.model_name,
306 "args": {
307 "job_id": job_id,
308 "img": img,
309 "num_characters": num_characters,
310 "zoom_factor": zoom_factor
311 }
312 }
315if __name__ == "__main__":
316 character_extractor = ImageCharacterExtractor()
318 img_filename = "generated_image_hidream_20250416T171929.png"
319 image = Image.open(img_filename).convert("RGB")
320 character_extractor.extract_characters(image, 2)