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

1import math 

2import cv2 

3import logging 

4import os 

5import numpy as np 

6 

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 

14 

15import torch 

16from torch import inference_mode 

17 

18from PIL import Image 

19 

20from wrapper_model import ModelGeneration 

21from image_utils import base64_to_img 

22 

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""" 

31 

32 

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) 

49 

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 

61 

62 

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 

73 

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)) 

77 

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) 

83 

84 # Center the crop on the box 

85 center_x = x + w // 2 

86 center_y = y + h // 2 

87 

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 

93 

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 

103 

104 # Final sanity check on box 

105 left = int(left) 

106 upper = int(upper) 

107 right = int(right) 

108 lower = int(lower) 

109 

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 

114 

115 

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] 

126 

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 

135 

136 

137class ImageCharacterExtractor(ModelGeneration): 

138 """ 

139 A class to zoom into a specific area of an image. 

140 """ 

141 

142 def __init__(self) -> None: 

143 super().__init__("yolo") 

144 

145 # Model components 

146 self.obj_recognition: Optional[YOLO] = None 

147 

148 def __del__(self) -> None: 

149 if self.obj_recognition is not None: 

150 self.obj_recognition = None 

151 

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") 

170 

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") 

181 

182 def init_model_parallelism(self) -> None: 

183 if self.world_size > 1: 

184 logging.warning("YOLO does not support distributed parallelism.") 

185 

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") 

195 

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) 

201 

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.") 

206 

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) 

216 

217 self._assert_model_init() 

218 assert self.obj_recognition is not None 

219 

220 self.running = True # We can run in parallel but good to know if we are running 

221 

222 try: 

223 results = self.obj_recognition.predict( 

224 img, 

225 verbose=False, 

226 device=self.device) 

227 

228 if len(results) == 0: 

229 logging.warning("No results for the image.") 

230 return [] 

231 

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) 

236 

237 if len(results[0].boxes) == 0: 

238 logging.warning("No characters detected in the image.") 

239 return [debug_img] + [None] * num_characters 

240 

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)) 

256 

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") 

265 

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) 

280 

281 def get_health(self) -> Dict[str, Any]: 

282 ret = super().get_health() 

283 ret["gpu"] = self.gpu 

284 return ret 

285 

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") 

292 

293 job_id = data_json.get("job_id", None) 

294 

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)) 

303 

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 } 

313 

314 

315if __name__ == "__main__": 

316 character_extractor = ImageCharacterExtractor() 

317 

318 img_filename = "generated_image_hidream_20250416T171929.png" 

319 image = Image.open(img_filename).convert("RGB") 

320 character_extractor.extract_characters(image, 2)