U
    di                     @   s   d dl Z d dlmZmZmZmZ ddlmZmZm	Z	m
Z
mZ ddlmZmZ e	 r^ddlmZ e r|d dlZddlmZmZ e
eZeeef Zee Zeed	d
G dd deZdS )    N)AnyDictListUnion   )add_end_docstringsis_torch_availableis_vision_availableloggingrequires_backends   )Pipelinebuild_pipeline_init_args)
load_image)(MODEL_FOR_OBJECT_DETECTION_MAPPING_NAMES,MODEL_FOR_TOKEN_CLASSIFICATION_MAPPING_NAMEST)Zhas_image_processorc                       sz   e Zd ZdZ fddZdd Zeeee	 f d fddZ
dd
dZdd ZdddZdeeef dddZ  ZS )ObjectDetectionPipelinea  
    Object detection pipeline using any `AutoModelForObjectDetection`. This pipeline predicts bounding boxes of objects
    and their classes.

    Example:

    ```python
    >>> from transformers import pipeline

    >>> detector = pipeline(model="facebook/detr-resnet-50")
    >>> detector("https://huggingface.co/datasets/Narsil/image_dummy/raw/main/parrots.png")
    [{'score': 0.997, 'label': 'bird', 'box': {'xmin': 69, 'ymin': 171, 'xmax': 396, 'ymax': 507}}, {'score': 0.999, 'label': 'bird', 'box': {'xmin': 398, 'ymin': 105, 'xmax': 767, 'ymax': 507}}]

    >>> # x, y  are expressed relative to the top left hand corner.
    ```

    Learn more about the basics of using a pipeline in the [pipeline tutorial](../pipeline_tutorial)

    This object detection pipeline can currently be loaded from [`pipeline`] using the following task identifier:
    `"object-detection"`.

    See the list of available models on [huggingface.co/models](https://huggingface.co/models?filter=object-detection).
    c                    sT   t  j|| | jdkr*td| j dt| d t }|t	 | 
| d S )NtfzThe z is only available in PyTorch.Zvision)super__init__	framework
ValueError	__class__r   r   copyupdater   Zcheck_model_type)selfargskwargsmappingr    K/tmp/pip-unpacked-wheel-bm_b0l5e/transformers/pipelines/object_detection.pyr   5   s    


z ObjectDetectionPipeline.__init__c                 K   sF   i }d|kr$t dt |d |d< i }d|kr<|d |d< |i |fS )NtimeoutzUThe `timeout` argument is deprecated and will be removed in version 5 of Transformers	threshold)warningswarnFutureWarning)r   r   Zpreprocess_paramsZpostprocess_kwargsr    r    r!   _sanitize_parameters@   s     z,ObjectDetectionPipeline._sanitize_parameters)returnc                    s,   d|krd|kr| d|d< t j||S )a  
        Detect objects (bounding boxes & classes) in the image(s) passed as inputs.

        Args:
            inputs (`str`, `List[str]`, `PIL.Image` or `List[PIL.Image]`):
                The pipeline handles three types of images:

                - A string containing an HTTP(S) link pointing to an image
                - A string containing a local path to an image
                - An image loaded in PIL directly

                The pipeline accepts either a single image or a batch of images. Images in a batch must all be in the
                same format: all as HTTP(S) links, all as local paths, or all as PIL images.
            threshold (`float`, *optional*, defaults to 0.5):
                The probability necessary to make a prediction.

        Return:
            A list of dictionaries or a list of list of dictionaries containing the result. If the input is a single
            image, will return a list of dictionaries, if the input is a list of several images, will return a list of
            list of dictionaries corresponding to each image.

            The dictionaries contain the following keys:

            - **label** (`str`) -- The class label identified by the model.
            - **score** (`float`) -- The score attributed by the model for that label.
            - **box** (`List[Dict[str, int]]`) -- The bounding box of detected object in image's original size.
        imagesinputs)popr   __call__)r   r   r   r   r    r!   r,   L   s    z ObjectDetectionPipeline.__call__Nc                 C   st   t ||d}t|j|jgg}| j|gdd}| jdkrF|| j}| j	d k	rh| j	|d |d dd}||d< |S )N)r"   pt)r)   return_tensorswordsboxes)textr0   r.   target_size)
r   torchZ	IntTensorheightwidthimage_processorr   toZtorch_dtype	tokenizer)r   imager"   r2   r*   r    r    r!   
preprocessm   s    

z"ObjectDetectionPipeline.preprocessc                 C   sB   | d}| jf |}|d|i|}| jd k	r>|d |d< |S )Nr2   bbox)r+   modelr   r8   )r   Zmodel_inputsr2   outputsmodel_outputsr    r    r!   _forwardx   s    

z ObjectDetectionPipeline._forward      ?c                    sN  |d }j d k	r|d  \  fdd|d djddjdd\}}fdd	| D }fd
d	|d dD }dddgfdd	t| ||D }nj||}	|	d }
|
d }|
d }|
d }| |
d< fdd	|D |
d< fdd	|D |
d< dddgfdd	t|
d |
d |
d D }|S )Nr2   r   c              
      sH    t| d  d  | d  d | d  d  | d  d gS )Nr   i  r   r      )_get_bounding_boxr3   ZTensor)r;   )r4   r   r5   r    r!   unnormalize   s    z8ObjectDetectionPipeline.postprocess.<locals>.unnormalizeZlogits)Zdimc                    s   g | ]} j jj| qS r    )r<   configid2label).0Z
predictionr   r    r!   
<listcomp>   s     z7ObjectDetectionPipeline.postprocess.<locals>.<listcomp>c                    s   g | ]} |qS r    r    )rG   r;   )rC   r    r!   rI      s     r;   Zscorelabelboxc                    s&   g | ]}|d  krt t |qS )r   dictziprG   vals)keysr#   r    r!   rI      s      scoreslabelsr0   c                    s   g | ]} j jj|  qS r    )r<   rE   rF   item)rG   rJ   rH   r    r!   rI      s     c                    s   g | ]}  |qS r    )rB   )rG   rK   rH   r    r!   rI      s     c                    s   g | ]}t t |qS r    rL   rO   )rQ   r    r!   rI      s   )r8   tolistZsqueezeZsoftmaxmaxrN   r6   Zpost_process_object_detection)r   r>   r#   r2   rR   classesrS   r0   
annotationZraw_annotationsZraw_annotationr    )r4   rQ   r   r#   rC   r5   r!   postprocess   s,    
"
"

z#ObjectDetectionPipeline.postprocessztorch.Tensor)rK   r(   c                 C   s8   | j dkrtd|  \}}}}||||d}|S )a%  
        Turns list [xmin, xmax, ymin, ymax] into dict { "xmin": xmin, ... }

        Args:
            box (`torch.Tensor`): Tensor containing the coordinates in corners format.

        Returns:
            bbox (`Dict[str, int]`): Dict containing the coordinates in corners format.
        r-   z9The ObjectDetectionPipeline is only available in PyTorch.)xminyminxmaxymax)r   r   intrU   )r   rK   rZ   r[   r\   r]   r;   r    r    r!   rB      s    

z)ObjectDetectionPipeline._get_bounding_box)N)r@   )__name__
__module____qualname____doc__r   r'   r   Predictionsr   
Predictionr,   r:   r?   rY   r   strr^   rB   __classcell__r    r    r   r!   r      s   !

-r   )r$   typingr   r   r   r   utilsr   r   r	   r
   r   baser   r   Zimage_utilsr   r3   Zmodels.auto.modeling_autor   r   Z
get_loggerr_   loggerre   rd   rc   r   r    r    r    r!   <module>   s   
