Aluode/PerceptionLabPortable
0
1# coding=utf-82# Copyright 2024 The HuggingFace Inc. team.3#4# Licensed under the Apache License, Version 2.0 (the "License");5# you may not use this file except in compliance with the License.6# You may obtain a copy of the License at7#8# http://www.apache.org/licenses/LICENSE-2.09#10# Unless required by applicable law or agreed to in writing, software11# distributed under the License is distributed on an "AS IS" BASIS,12# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.13# See the License for the specific language governing permissions and14# limitations under the License.15"""16Processor class for Grounding DINO.17"""18 19import pathlib20import warnings21from typing import TYPE_CHECKING, Optional, Union22 23from ...image_transforms import center_to_corners_format24from ...image_utils import AnnotationFormat, ImageInput25from ...processing_utils import ImagesKwargs, ProcessingKwargs, ProcessorMixin, Unpack26from ...tokenization_utils_base import BatchEncoding, PreTokenizedInput, TextInput27from ...utils import TensorType, is_torch_available28 29 30if is_torch_available():31 import torch32 33if TYPE_CHECKING:34 from .modeling_grounding_dino import GroundingDinoObjectDetectionOutput35 36 37AnnotationType = dict[str, Union[int, str, list[dict]]]38 39 40def get_phrases_from_posmap(posmaps, input_ids):41 """Get token ids of phrases from posmaps and input_ids.42 43 Args:44 posmaps (`torch.BoolTensor` of shape `(num_boxes, hidden_size)`):45 A boolean tensor of text-thresholded logits related to the detected bounding boxes.46 input_ids (`torch.LongTensor`) of shape `(sequence_length, )`):47 A tensor of token ids.48 """49 left_idx = 050 right_idx = posmaps.shape[-1] - 151 52 # Avoiding altering the input tensor53 posmaps = posmaps.clone()54 55 posmaps[:, 0 : left_idx + 1] = False56 posmaps[:, right_idx:] = False57 58 token_ids = []59 for posmap in posmaps:60 non_zero_idx = posmap.nonzero(as_tuple=True)[0].tolist()61 token_ids.append([input_ids[i] for i in non_zero_idx])62 63 return token_ids64 65 66def _is_list_of_candidate_labels(text) -> bool:67 """Check that text is list/tuple of strings and each string is a candidate label and not merged candidate labels text.68 Merged candidate labels text is a string with candidate labels separated by a dot.69 """70 if isinstance(text, (list, tuple)):71 return all(isinstance(t, str) and "." not in t for t in text)72 return False73 74 75def _merge_candidate_labels_text(text: list[str]) -> str:76 """77 Merge candidate labels text into a single string. Ensure all labels are lowercase.78 For example, ["A cat", "a dog"] -> "a cat. a dog."79 """80 labels = [t.strip().lower() for t in text] # ensure lowercase81 merged_labels_str = ". ".join(labels) + "." # join with dot and add a dot at the end82 return merged_labels_str83 84 85class DictWithDeprecationWarning(dict):86 message = (87 "The key `labels` is will return integer ids in `GroundingDinoProcessor.post_process_grounded_object_detection` "88 "output since v4.51.0. Use `text_labels` instead to retrieve string object names."89 )90 91 def __getitem__(self, key):92 if key == "labels":93 warnings.warn(self.message, FutureWarning)94 return super().__getitem__(key)95 96 def get(self, key, *args, **kwargs):97 if key == "labels":98 warnings.warn(self.message, FutureWarning)99 return super().get(key, *args, **kwargs)100 101 102class GroundingDinoImagesKwargs(ImagesKwargs, total=False):103 annotations: Optional[Union[AnnotationType, list[AnnotationType]]]104 return_segmentation_masks: Optional[bool]105 masks_path: Optional[Union[str, pathlib.Path]]106 do_convert_annotations: Optional[bool]107 format: Optional[Union[str, AnnotationFormat]]108 109 110class GroundingDinoProcessorKwargs(ProcessingKwargs, total=False):111 images_kwargs: GroundingDinoImagesKwargs112 _defaults = {113 "text_kwargs": {114 "add_special_tokens": True,115 "padding": False,116 "stride": 0,117 "return_overflowing_tokens": False,118 "return_special_tokens_mask": False,119 "return_offsets_mapping": False,120 "return_token_type_ids": True,121 "return_length": False,122 "verbose": True,123 }124 }125 126 127class GroundingDinoProcessor(ProcessorMixin):128 r"""129 Constructs a Grounding DINO processor which wraps a Deformable DETR image processor and a BERT tokenizer into a130 single processor.131 132 [`GroundingDinoProcessor`] offers all the functionalities of [`GroundingDinoImageProcessor`] and133 [`AutoTokenizer`]. See the docstring of [`~GroundingDinoProcessor.__call__`] and [`~GroundingDinoProcessor.decode`]134 for more information.135 136 Args:137 image_processor (`GroundingDinoImageProcessor`):138 An instance of [`GroundingDinoImageProcessor`]. The image processor is a required input.139 tokenizer (`AutoTokenizer`):140 An instance of ['PreTrainedTokenizer`]. The tokenizer is a required input.141 """142 143 attributes = ["image_processor", "tokenizer"]144 image_processor_class = "GroundingDinoImageProcessor"145 tokenizer_class = "AutoTokenizer"146 valid_processor_kwargs = GroundingDinoProcessorKwargs147 148 def __init__(self, image_processor, tokenizer):149 super().__init__(image_processor, tokenizer)150 151 def __call__(152 self,153 images: Optional[ImageInput] = None,154 text: Union[TextInput, PreTokenizedInput, list[TextInput], list[PreTokenizedInput]] = None,155 **kwargs: Unpack[GroundingDinoProcessorKwargs],156 ) -> BatchEncoding:157 """158 This method uses [`GroundingDinoImageProcessor.__call__`] method to prepare image(s) for the model, and159 [`BertTokenizerFast.__call__`] to prepare text for the model.160 161 Args:162 images (`ImageInput`, `list[ImageInput]`, *optional*):163 The image or batch of images to be processed. The image might be either PIL image, numpy array or a torch tensor.164 text (`TextInput`, `PreTokenizedInput`, `list[TextInput]`, `list[PreTokenizedInput]`, *optional*):165 Candidate labels to be detected on the image. The text might be one of the following:166 - A list of candidate labels (strings) to be detected on the image (e.g. ["a cat", "a dog"]).167 - A batch of candidate labels to be detected on the batch of images (e.g. [["a cat", "a dog"], ["a car", "a person"]]).168 - A merged candidate labels string to be detected on the image, separated by "." (e.g. "a cat. a dog.").169 - A batch of merged candidate labels text to be detected on the batch of images (e.g. ["a cat. a dog.", "a car. a person."]).170 """171 if text is not None:172 text = self._preprocess_input_text(text)173 return super().__call__(images=images, text=text, **kwargs)174 175 def _preprocess_input_text(self, text):176 """177 Preprocess input text to ensure that labels are in the correct format for the model.178 If the text is a list of candidate labels, merge the candidate labels into a single string,179 for example, ["a cat", "a dog"] -> "a cat. a dog.". In case candidate labels are already in a form of180 "a cat. a dog.", the text is returned as is.181 """182 183 if _is_list_of_candidate_labels(text):184 text = _merge_candidate_labels_text(text)185 186 # for batched input187 elif isinstance(text, (list, tuple)) and all(_is_list_of_candidate_labels(t) for t in text):188 text = [_merge_candidate_labels_text(sample) for sample in text]189 190 return text191 192 def post_process_grounded_object_detection(193 self,194 outputs: "GroundingDinoObjectDetectionOutput",195 input_ids: Optional[TensorType] = None,196 threshold: float = 0.25,197 text_threshold: float = 0.25,198 target_sizes: Optional[Union[TensorType, list[tuple]]] = None,199 text_labels: Optional[list[list[str]]] = None,200 ):201 """202 Converts the raw output of [`GroundingDinoForObjectDetection`] into final bounding boxes in (top_left_x, top_left_y,203 bottom_right_x, bottom_right_y) format and get the associated text label.204 205 Args:206 outputs ([`GroundingDinoObjectDetectionOutput`]):207 Raw outputs of the model.208 input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):209 The token ids of the input text. If not provided will be taken from the model output.210 threshold (`float`, *optional*, defaults to 0.25):211 Threshold to keep object detection predictions based on confidence score.212 text_threshold (`float`, *optional*, defaults to 0.25):213 Score threshold to keep text detection predictions.214 target_sizes (`torch.Tensor` or `list[tuple[int, int]]`, *optional*):215 Tensor of shape `(batch_size, 2)` or list of tuples (`tuple[int, int]`) containing the target size216 `(height, width)` of each image in the batch. If unset, predictions will not be resized.217 text_labels (`list[list[str]]`, *optional*):218 List of candidate labels to be detected on each image. At the moment it's *NOT used*, but required219 to be in signature for the zero-shot object detection pipeline. Text labels are instead extracted220 from the `input_ids` tensor provided in `outputs`.221 222 Returns:223 `list[Dict]`: A list of dictionaries, each dictionary containing the224 - **scores**: tensor of confidence scores for detected objects225 - **boxes**: tensor of bounding boxes in [x0, y0, x1, y1] format226 - **labels**: list of text labels for each detected object (will be replaced with integer ids in v4.51.0)227 - **text_labels**: list of text labels for detected objects228 """229 batch_logits, batch_boxes = outputs.logits, outputs.pred_boxes230 input_ids = input_ids if input_ids is not None else outputs.input_ids231 232 if target_sizes is not None and len(target_sizes) != len(batch_logits):233 raise ValueError("Make sure that you pass in as many target sizes as the batch dimension of the logits")234 235 batch_probs = torch.sigmoid(batch_logits) # (batch_size, num_queries, 256)236 batch_scores = torch.max(batch_probs, dim=-1)[0] # (batch_size, num_queries)237 238 # Convert to [x0, y0, x1, y1] format239 batch_boxes = center_to_corners_format(batch_boxes)240 241 # Convert from relative [0, 1] to absolute [0, height] coordinates242 if target_sizes is not None:243 if isinstance(target_sizes, list):244 img_h = torch.Tensor([i[0] for i in target_sizes])245 img_w = torch.Tensor([i[1] for i in target_sizes])246 else:247 img_h, img_w = target_sizes.unbind(1)248 249 scale_fct = torch.stack([img_w, img_h, img_w, img_h], dim=1).to(batch_boxes.device)250 batch_boxes = batch_boxes * scale_fct[:, None, :]251 252 results = []253 for idx, (scores, boxes, probs) in enumerate(zip(batch_scores, batch_boxes, batch_probs)):254 keep = scores > threshold255 scores = scores[keep]256 boxes = boxes[keep]257 258 # extract text labels259 prob = probs[keep]260 label_ids = get_phrases_from_posmap(prob > text_threshold, input_ids[idx])261 objects_text_labels = self.batch_decode(label_ids)262 263 result = DictWithDeprecationWarning(264 {265 "scores": scores,266 "boxes": boxes,267 "text_labels": objects_text_labels,268 # TODO: @pavel, set labels to None since v4.51.0 or find a way to extract ids269 "labels": objects_text_labels,270 }271 )272 results.append(result)273 274 return results275 276 277__all__ = ["GroundingDinoProcessor"]278 