
    ^j)                         d Z ddlmZ ddlZddlmc mZ ddlmZ	 ddl
mZ ddlmZ ddlmZ dd	lmZmZ dd
lmZmZmZmZmZmZ ddlmZmZ ddlmZmZm Z   G d ded      Z!e G d de             Z"dgZ#y)zImage processor class for BEiT.    )UnionN)
functional   )TorchvisionBackend)'SemanticSegmentationPostProcessorOutput)BatchFeature)group_images_by_shapereorder_images)IMAGENET_STANDARD_MEANIMAGENET_STANDARD_STDChannelDimension
ImageInputPILImageResamplingSizeDict)ImagesKwargsUnpack)
TensorTypeauto_docstringis_torch_availablec                       e Zd ZU dZeed<   y)BeitImageProcessorKwargsak  
    do_reduce_labels (`bool`, *optional*, defaults to `self.do_reduce_labels`):
        Whether or not to reduce all label values of segmentation maps by 1. Usually used for datasets where 0
        is used for background, and background itself is not included in all classes of a dataset (e.g.
        ADE20k). The background label will be replaced by 255.
    do_reduce_labelsN)__name__
__module____qualname____doc__bool__annotations__     y/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/transformers/models/beit/image_processing_beit.pyr   r   &   s     r    r   F)totalc                       e Zd ZdZeZej                  Ze	Z
eZdddZdZdddZdZdZdZdZdZdee   f fdZe	 d(d	ed
edz  dee   def fd       Z	 d(d	ed
edz  dededeez  dz  deedf   dz  defdZ de!d   de!d   fdZ"	 d)d	e!d   dede#dddede#dede$dede$e!e$   z  dz  d e$e!e$   z  dz  d!edz  d"ede!d   fd#Z%	 d*d$e!e&   dz  d%edd&fd'Z' xZ(S )+BeitImageProcessorz/PIL backend for BEiT with reduce_label support.   )heightwidthTFkwargsc                 $    t        |   di | y )Nr   )super__init__)selfr(   	__class__s     r!   r+   zBeitImageProcessor.__init__C   s    "6"r    Nimagessegmentation_mapsreturnc                 &    t        |   ||fi |S )zp
        segmentation_maps (`ImageInput`, *optional*):
            The segmentation maps to preprocess.
        )r*   
preprocess)r,   r.   r/   r(   r-   s       r!   r2   zBeitImageProcessor.preprocessF   s     w!&*;FvFFr    do_convert_rgbinput_data_formatreturn_tensorsdeviceztorch.devicec                    | j                  ||||      }|j                         }d|d<   i }	 | j                  |fi ||	d<   || j                  |ddt        j                        }
|j                         }|j                  ddd        | j                  dd|
i|}
|
D cg c]0  }|j                  d	      j                  t        j                        2 }
}|
|	d
<   t        |	|      S c c}w )z"Handle extra inputs beyond images.)r.   r3   r4   r6   Fr   pixel_values   )r.   expected_ndimsr3   r4   )do_normalize
do_rescaler.   r   labels)datatensor_typer   )_prepare_image_like_inputscopy_preprocessr   FIRSTupdatesqueezetotorchint64r   )r,   r.   r/   r3   r4   r5   r6   r(   images_kwargsr>   processed_segmentation_mapssegmentation_maps_kwargsprocessed_segmentation_maps                r!   _preprocess_image_like_inputsz0BeitImageProcessor._preprocess_image_like_inputsS   s+    00.L]fl 1 
 ,1()/t//H-H^ (*.*I*I( $"2"8"8	 +J +' (.{{}$$++URW,XY*:$*:*: +2+6N+' 3N+. +221588E+' + 9DN>BB+s   $5C-r=   ztorch.Tensorc           	      f   t        t        |            D ]  }||   }t        j                  |dk(  t        j                  d|j
                  |j                        |      }|dz
  }t        j                  |dk(  t        j                  d|j
                  |j                        |      }|||<    |S )z/Reduce label values by 1, replacing 0 with 255.r      )dtyper6         )rangelenrG   wheretensorrP   r6   )r,   r=   idxlabels       r!   reduce_labelzBeitImageProcessor.reduce_label   s    V% 	 C3KEKK
ELLEKKX]XdXd,eglmEAIEKKell3ekkZ_ZfZf.ginoEF3K	  r    	do_resizesizeresamplez7PILImageResampling | tvF.InterpolationMode | int | Nonedo_center_crop	crop_sizer<   rescale_factorr;   
image_mean	image_stddisable_groupingr   c           	         |r| j                  |      }t        ||      \  }}i }|j                         D ]  \  }}|r| j                  |||      }|||<   ! t	        ||      }t        ||      \  }}i }|j                         D ]4  \  }}|r| j                  ||      }| j                  ||||	|
|      }|||<   6 t	        ||      }|S )zCustom preprocessing for BEiT.)rb   )rY   r	   itemsresizer
   center_croprescale_and_normalize)r,   r.   rZ   r[   r\   r]   r^   r<   r_   r;   r`   ra   rb   r   r(   grouped_imagesgrouped_images_indexresized_images_groupedshapestacked_imagesresized_imagesprocessed_images_groupedprocessed_imagess                          r!   rB   zBeitImageProcessor._preprocess   s   $ &&v.F 0EV^n/o,,!#%3%9%9%; 	;!E>!%^T8!L,:"5)	; ((>@TU 0E^fv/w,,#% %3%9%9%; 	=!E>!%!1!1.)!L!77
NL*V_N /=$U+	= **BDXYr    target_sizesreturn_segmentation_scoreszBlist[torch.Tensor] | list[SemanticSegmentationPostProcessorOutput]c                    t               st        d      |j                  }|t        |      t        |      k7  rt	        d      t        |t        j                        r|j                         }g }t        t        |            D ]g  }t        j                  ||   j                  d      ||   dd      }|d   j                  d      }|j                  t        ||d   d	             i nJ|j                  d
      }	t        |j                   d         D 
cg c]  }
t        |	|
   ||
   d	       }}
|s|D cg c]  }|j"                   }}|S c c}
w c c}w )a  
        Converts the output of [`BeitForSemanticSegmentation`] into semantic segmentation maps.

        Args:
            outputs ([`BeitForSemanticSegmentation`]):
                Raw outputs of the model.
            target_sizes (`list[Tuple]` of length `batch_size`, *optional*):
                List of tuples corresponding to the requested final size (height, width) of each prediction. If unset,
                predictions will not be resized.
            return_segmentation_scores (`bool`, *optional*, defaults to `False`):
                Whether to return segmentation scores alongside the segmentation map. When `True`, each element of
                the returned list is a [`SemanticSegmentationPostProcessorOutput`] with fields `segmentation`
                (class IDs, shape `(height, width)`) and `segmentation_scores` (shape `(num_classes, height, width)`).

        Returns:
            `list[torch.Tensor]` or `list[SemanticSegmentationPostProcessorOutput]`: When
            `return_segmentation_scores=False` (default), a list of length `batch_size` where each item is a
            segmentation map of shape `(height, width)` with class IDs. When `return_segmentation_scores=True`,
            a list of [`SemanticSegmentationPostProcessorOutput`] with fields `segmentation` (class IDs, shape
            `(height, width)`) and `segmentation_scores` (shape `(num_classes, height, width)`). In both cases,
            `(height, width)` corresponds to the target size (if `target_sizes` is specified).
        z:PyTorch is required for post_process_semantic_segmentationzTMake sure that you pass in as many target sizes as the batch dimension of the logitsr   )dimbilinearF)r[   modealign_corners)segmentationsegmentation_scores)r>   rQ   )r   ImportErrorlogitsrT   
ValueError
isinstancerG   TensornumpyrS   Finterpolate	unsqueezeargmaxappendr   rk   rw   )r,   outputsrp   rq   rz   semantic_segmentationrW   resized_logitssemantic_mapseg_mapsiitems               r!   "post_process_semantic_segmentationz5BeitImageProcessor.post_process_semantic_segmentation   sy   2 "#Z[[ #6{c,// j  ,5+113$&!S[) 	!"3K))a)0|C7Hzin"  .a077A7>%,,;.:SabcSde	 }}}+H
 v||A/	%  8*21+fUViX%! % *CX$Y4T%6%6$Y!$Y$$% %Zs   EE)N)F)NF))r   r   r   r   r   valid_kwargsr   BICUBICr\   r   r`   r   ra   r[   default_to_squarer^   rZ   r]   r<   r;   r   r   r+   r   r   r   r2   r   r   strr   r   rM   listrY   r   floatrB   tupler   __classcell__)r-   s   @r!   r$   r$   1   s9   9+L!))H'J%IC(D-IINJL#(@!A #  04
G
G &,
G 12	
G
 

G 
G& 59*C*C &,*C 	*C
 ,*C j(4/*C c>)*T1*C 
*CX4#7 D<P 0 "', ^$,  ,  	, 
 L,  ,  ,  ,  ,  ,  DK'$.,  4;&-,  +,  ,   
n	!, ^ di@%%)%[4%7@%\`@%	M@%r    r$   )$r   typingr   rG   torch.nn.functionalnnr   r   torchvision.transforms.v2tvFimage_processing_backendsr   image_processing_outputsr   image_processing_utilsr   image_transformsr	   r
   image_utilsr   r   r   r   r   r   processing_utilsr   r   utilsr   r   r   r   r$   __all__r   r    r!   <module>r      sr    &     7 ; O 2 E  5 C C|5  E%+ E% E%P  
 r    