diffusers/Florence2-image-Annotator
120
1from typing import List, Union2 3import numpy as np4import torch5from diffusers.modular_pipelines import (6 ComponentSpec,7 InputParam,8 ModularPipelineBlocks,9 OutputParam,10 PipelineState,11)12from PIL import Image, ImageDraw13from transformers import AutoProcessor, Florence2ForConditionalGeneration14 15 16class Florence2ImageAnnotatorBlock(ModularPipelineBlocks):17 @property18 def expected_components(self):19 return [20 ComponentSpec(21 name="image_annotator",22 type_hint=Florence2ForConditionalGeneration,23 repo="florence-community/Florence-2-base-ft",24 ),25 ComponentSpec(26 name="image_annotator_processor",27 type_hint=AutoProcessor,28 repo="florence-community/Florence-2-base-ft",29 ),30 ]31 32 @property33 def inputs(self) -> List[InputParam]:34 return [35 InputParam(36 "image",37 type_hint=Union[Image.Image, List[Image.Image]],38 required=True,39 description="Image(s) to annotate",40 metadata={"mellon":"image"},41 ),42 InputParam(43 "annotation_task",44 type_hint=Union[str, List[str]],45 default="<REFERRING_EXPRESSION_SEGMENTATION>",46 metadata={"mellon":"dropdown"},47 description="""Annotation Task to perform on the image.48 Supported Tasks:49 50 <OD>51 <REFERRING_EXPRESSION_SEGMENTATION>52 <CAPTION>53 <DETAILED_CAPTION>54 <MORE_DETAILED_CAPTION>55 <DENSE_REGION_CAPTION>56 <REGION_PROPOSAL>57 <CAPTION_TO_PHRASE_GROUNDING>58 <OPEN_VOCABULARY_DETECTION>59 <OCR>60 <OCR_WITH_REGION>61 62 """,63 ),64 InputParam(65 "annotation_prompt",66 type_hint=Union[str, List[str]],67 required=True,68 metadata={"mellon":"textbox"},69 description="""Annotation Prompt to provide more context to the task.70 Can be used to detect or segment out specific elements in the image71 """,72 ),73 InputParam(74 "annotation_output_type",75 type_hint=str,76 default="mask_image",77 metadata={"mellon":"dropdown"},78 description="""Output type from annotation predictions. Availabe options are79 annotation:80 - raw annotation predictions from the model based on task type.81 mask_image:82 -black and white mask image for the given image based on the task type83 mask_overlay:84 - white mask overlayed on the original image85 bounding_box:86 - bounding boxes drawn on the original image87 """,88 ),89 InputParam(90 "annotation_overlay",91 type_hint=bool,92 required=True,93 default=False,94 description="",95 metadata={"mellon":"checkbox"},96 ),97 InputParam(98 "fill",99 type_hint=str,100 default="white",101 description="",102 ),103 ]104 105 @property106 def intermediate_outputs(self) -> List[OutputParam]:107 return [108 OutputParam(109 "annotations",110 type_hint=dict,111 description="Annotations Predictions for input Image(s)",112 ),113 OutputParam(114 "images",115 type_hint=Image,116 description="Annotated input Image(s)",117 metadata={"mellon":"image"},118 ),119 ]120 121 def get_annotations(self, components, images, prompts, task):122 task_prompts = [task + prompt for prompt in prompts]123 124 inputs = components.image_annotator_processor(125 text=task_prompts, images=images, return_tensors="pt"126 ).to(components.image_annotator.device, components.image_annotator.dtype)127 128 generated_ids = components.image_annotator.generate(129 input_ids=inputs["input_ids"],130 pixel_values=inputs["pixel_values"],131 max_new_tokens=1024,132 early_stopping=False,133 do_sample=False,134 num_beams=3,135 )136 annotations = components.image_annotator_processor.batch_decode(137 generated_ids, skip_special_tokens=False138 )139 140 outputs = []141 for image, annotation in zip(images, annotations):142 outputs.append(143 components.image_annotator_processor.post_process_generation(144 annotation, task=task, image_size=(image.width, image.height)145 )146 )147 148 return outputs149 150 def _iter_polygon_point_sets(self, poly):151 """152 Yields lists of (x, y) points for all simple polygons found in `poly`.153 Supports formats:154 - [x1, y1, x2, y2, ...]155 - [[x, y], [x, y], ...]156 - [xs, ys]157 - dict {'x': xs, 'y': ys}158 - nested lists containing any of the above159 """160 if poly is None:161 return162 163 def is_num(v):164 return isinstance(v, (int, float, np.number))165 166 # dict {'x': [...], 'y': [...]}167 if isinstance(poly, dict) and "x" in poly and "y" in poly:168 xs, ys = poly["x"], poly["y"]169 if (170 isinstance(xs, (list, tuple))171 and isinstance(ys, (list, tuple))172 and len(xs) == len(ys)173 ):174 pts = list(zip(xs, ys))175 if len(pts) >= 3:176 yield pts177 return178 179 if isinstance(poly, (list, tuple)):180 # flat numeric [x1, y1, ...]181 if all(is_num(v) for v in poly):182 coords = list(poly)183 if len(coords) >= 6 and len(coords) % 2 == 0:184 yield list(zip(coords[0::2], coords[1::2]))185 return186 187 # list of pairs [[x, y], ...]188 if all(189 isinstance(v, (list, tuple))190 and len(v) == 2191 and all(is_num(n) for n in v)192 for v in poly193 ):194 if len(poly) >= 3:195 yield [tuple(v) for v in poly]196 return197 198 # [xs, ys]199 if len(poly) == 2 and all(isinstance(v, (list, tuple)) for v in poly):200 xs, ys = poly201 try:202 if len(xs) == len(ys) and len(xs) >= 3:203 yield list(zip(xs, ys))204 return205 except TypeError:206 pass207 208 # nested: recurse into parts209 for part in poly:210 yield from self._iter_polygon_point_sets(part)211 # other types are ignored212 213 def prepare_mask(self, images, annotations, overlay=False, fill="white"):214 masks = []215 for image, annotation in zip(images, annotations):216 mask_image = image.copy() if overlay else Image.new("L", image.size, 0)217 draw = ImageDraw.Draw(mask_image)218 219 # use a safe fill for grayscale masks220 mask_fill = fill221 if not overlay and isinstance(fill, str):222 # for "L" mode, white -> 255223 mask_fill = 255224 225 for _, _annotation in annotation.items():226 if "polygons" in _annotation:227 for poly in _annotation["polygons"]:228 for pts in self._iter_polygon_point_sets(poly):229 if len(pts) < 3:230 continue231 # clip to image bounds and flatten232 flat = []233 for x, y in pts:234 xi = int(round(max(0, min(image.width - 1, x))))235 yi = int(round(max(0, min(image.height - 1, y))))236 flat.extend([xi, yi])237 draw.polygon(flat, fill=mask_fill)238 239 elif "bboxes" in _annotation:240 for bbox in _annotation["bboxes"]:241 flat = np.array(bbox).flatten().tolist()242 if len(flat) == 4:243 x0, y0, x1, y1 = flat244 draw.rectangle(245 (246 int(round(x0)),247 int(round(y0)),248 int(round(x1)),249 int(round(y1)),250 ),251 fill=mask_fill,252 )253 254 elif "quad_boxes" in _annotation:255 for quad in _annotation["quad_boxes"]:256 for pts in self._iter_polygon_point_sets(quad):257 if len(pts) < 3:258 continue259 flat = []260 for x, y in pts:261 xi = int(round(max(0, min(image.width - 1, x))))262 yi = int(round(max(0, min(image.height - 1, y))))263 flat.extend([xi, yi])264 draw.polygon(flat, fill=mask_fill)265 266 masks.append(mask_image)267 268 return masks269 270 def prepare_bounding_boxes(self, images, annotations):271 outputs = []272 for image, annotation in zip(images, annotations):273 image_copy = image.copy()274 draw = ImageDraw.Draw(image_copy)275 for _, _annotation in annotation.items():276 # Standard axis-aligned boxes277 bboxes = _annotation.get("bboxes", [])278 labels = _annotation.get("labels", [])279 280 if len(labels) == 0:281 labels = _annotation.get("bboxes_labels", [])282 283 for i, bbox in enumerate(bboxes):284 flat = np.array(bbox).flatten().tolist()285 286 if len(flat) != 4:287 continue288 289 x0, y0, x1, y1 = flat290 draw.rectangle(291 (292 int(round(x0)),293 int(round(y0)),294 int(round(x1)),295 int(round(y1)),296 ),297 outline="red",298 width=3,299 )300 label = labels[i] if i < len(labels) else ""301 if label:302 text_y = max(0, int(y0) - 20)303 draw.text((int(x0), text_y), label, fill="red")304 305 # Quadrilateral boxes (draw as polygons)306 quad_boxes = _annotation.get("quad_boxes", [])307 qlabels = _annotation.get("labels", [])308 for i, quad in enumerate(quad_boxes):309 for pts in self._iter_polygon_point_sets(quad):310 if len(pts) < 3:311 continue312 flat = []313 xs, ys = [], []314 for x, y in pts:315 xi = int(round(max(0, min(image.width - 1, x))))316 yi = int(round(max(0, min(image.height - 1, y))))317 flat.extend([xi, yi])318 xs.append(xi)319 ys.append(yi)320 321 # Outline polygon322 try:323 draw.polygon(flat, outline="red", width=3)324 except TypeError:325 # Pillow without width for polygon326 draw.polygon(flat, outline="red")327 328 # Optional label at centroid (inside the quad)329 label = qlabels[i] if i < len(qlabels) else ""330 if label:331 cx = int(round(sum(xs) / len(xs)))332 cy = int(round(sum(ys) / len(ys)))333 cx = max(0, min(image.width - 1, cx))334 cy = max(0, min(image.height - 1, cy))335 draw.text((cx, cy), label, fill="red")336 337 outputs.append(image_copy)338 339 return outputs340 341 def prepare_inputs(self, images, prompts):342 prompts = prompts or ""343 344 if isinstance(images, Image.Image):345 images = [images]346 if isinstance(prompts, str):347 prompts = [prompts]348 349 if len(images) != len(prompts):350 raise ValueError("Number of images and annotation prompts must match.")351 352 return images, prompts353 354 @torch.no_grad()355 def __call__(self, components, state: PipelineState) -> PipelineState:356 block_state = self.get_block_state(state)357 skip_image = False358 359 # these don't require a prompt and fail if one is given360 if (361 block_state.annotation_task == "<OD>"362 or block_state.annotation_task == "<DENSE_REGION_CAPTION>"363 or block_state.annotation_task == "<REGION_PROPOSAL>"364 or block_state.annotation_task == "<OCR_WITH_REGION>"365 ):366 block_state.annotation_prompt = ""367 block_state.annotation_output_type = "bounding_box"368 # these don't require a prompt and doesn't ouput an image369 elif (370 block_state.annotation_task == "<CAPTION>"371 or block_state.annotation_task == "<DETAILED_CAPTION>"372 or block_state.annotation_task == "<MORE_DETAILED_CAPTION>"373 or block_state.annotation_task == "<OCR>"374 ):375 block_state.annotation_prompt = ""376 skip_image = True377 378 images, annotation_task_prompt = self.prepare_inputs(379 block_state.image, block_state.annotation_prompt380 )381 task = block_state.annotation_task382 fill = block_state.fill383 384 annotations = self.get_annotations(385 components, images, annotation_task_prompt, task386 )387 388 block_state.annotations = annotations389 block_state.images = None390 391 if not skip_image:392 if block_state.annotation_output_type == "mask_image":393 block_state.images = self.prepare_mask(images, annotations)394 395 if block_state.annotation_output_type == "mask_overlay":396 block_state.images = self.prepare_mask(397 images, annotations, overlay=True, fill=fill398 )399 elif block_state.annotation_output_type == "bounding_box":400 block_state.images = self.prepare_bounding_boxes(images, annotations)401 402 self.set_block_state(state, block_state)403 404 return components, state405 