Team Ai
Modelpublic

diffusers/Florence2-image-Annotator

sourceHugging Faceupdated 8mo agoView on Hugging Face
1likes20downloads
block.py405 linesDownload Raw Back to root
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