
    ^j2                     8   d Z ddlmZ ddl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mZmZ dd	lmZmZ dd
lmZmZmZ erddlmZ  e       rddlmZmZ ddl m!Z" d Z#defdZ$	 	 ddZ%	 	 	 	 ddZ& G d ded      Z'e G d de             Z(dgZ)y)z$Image processor class for SuperGlue.    )TYPE_CHECKINGN   )TorchvisionBackend)BatchFeature)group_images_by_shapereorder_images)
ImageInput	ImageTypePILImageResamplingSizeDictget_image_typeis_pil_imageis_valid_imageto_numpy_array)ImagesKwargsUnpack)
TensorTypeauto_docstringis_vision_available   )SuperGlueKeypointMatchingOutput)Image	ImageDraw)
functionalc                     t        |       xsC t        |       xr6 t        |       t        j                  k7  xr t        | j                        dk(  S )Nr   r   r   r   r
   PILlenshapeimages    /var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/transformers/models/superglue/image_processing_superglue.py_is_valid_imager#   2   sF     ub."79=="HbSQVQ\Q\M]abMb    imagesc                     d}d t        | t              rQt        |       dk(  rt        fd| D              r| S t        fd| D              r| D cg c]  }|D ]  }|  c}}S t	        |      c c}}w )N)z-Input images must be a one of the following :z - A pair of PIL images.z - A pair of 3D arrays.z! - A list of pairs of PIL images.z  - A list of pairs of 3D arrays.c                     t        |       xsC t        |       xr6 t        |       t        j                  k7  xr t        | j                        dk(  S )z$images is a PIL Image or a 3D array.r   r   r    s    r"   r#   z8validate_and_format_image_pairs.<locals>._is_valid_imageA   sG    E" 
5!fnU&;y}}&LfQTUZU`U`QaefQf	
r$      c              3   .   K   | ]  } |        y wN .0r!   r#   s     r"   	<genexpr>z2validate_and_format_image_pairs.<locals>.<genexpr>H   s     #Q_U%;#Q   c              3      K   | ]:  }t        |t              xr$ t        |      d k(  xr t        fd|D               < yw)r(   c              3   .   K   | ]  } |        y wr*   r+   r,   s     r"   r.   z<validate_and_format_image_pairs.<locals>.<genexpr>.<genexpr>M   s     CuOE*Cr/   N)
isinstancelistr   all)r-   
image_pairr#   s     r"   r.   z2validate_and_format_image_pairs.<locals>.<genexpr>J   sN      
  z4( DJ1$DC
CCD
s   A A)r2   r3   r   r4   
ValueError)r%   error_messager5   r!   r#   s       @r"   validate_and_format_image_pairsr8   8   s    M
 &$v;!#Q&#Q QM 
 %	
 
 -3Kj
KuEKEKK
]
## Ls   A3c           	      $   | j                   dk  s#| j                  | j                   dk(  rdnd   dk(  ryt        j                  | ddddddf   | ddddddf   k(        xr. t        j                  | ddddddf   | ddddddf   k(        S )zAChecks if an image is grayscale (all RGB channels are identical).r   r   r   T.Nr(   )ndimr   torchr4   r    s    r"   is_grayscaler<   T   s     zzA~%**/QqAQF99U31a<(E#q!Q,,??@ UYYc1aluS!Q\22F r$   c                 J    t        |       r| S t        j                  | d      S )a  
    Converts an image to grayscale format using the NTSC formula. Only support torch.Tensor.

    This function is supposed to return a 1-channel image, but it returns a 3-channel image with the same value in each
    channel, because of an issue that is discussed in :
    https://github.com/huggingface/transformers/pull/25786#issuecomment-1730176446

    Args:
        image (torch.Tensor):
            The image to convert.
    r   )num_output_channels)r<   tvFrgb_to_grayscaler    s    r"   convert_to_grayscalerA   _   s$     E1==r$   c                       e Zd ZU dZeed<   y)SuperGlueImageProcessorKwargsz
    do_grayscale (`bool`, *optional*, defaults to `self.do_grayscale`):
        Whether to convert the image to grayscale. Can be overridden by `do_grayscale` in the `preprocess` method.
    do_grayscaleN)__name__
__module____qualname____doc__bool__annotations__r+   r$   r"   rC   rC   r   s    
 r$   rC   F)totalc                   x    e Zd ZeZej                  ZdddZdZ	dZ
dZdZdZdZdee   f fd	Zed
edee   def fd       Zd
edefdZ	 d"d
ed   dededddedededz  deez  dz  dedefdZ	 d#dddeee   z  dedeeeej@                  f      fdZ!deeeej@                  f      ded   fd Z"d! Z# xZ$S )$SuperGlueImageProcessori  i  )heightwidthFTgp?Nkwargsc                 $    t        |   di | y )Nr+   )super__init__)selfrP   	__class__s     r"   rS   z SuperGlueImageProcessor.__init__   s    "6"r$   r%   returnc                 $    t        |   |fi |S r*   )rR   
preprocess)rT   r%   rP   rU   s      r"   rX   z"SuperGlueImageProcessor.preprocess   s    w!&3F33r$   c                 :    | j                  |      }t        |      S r*   )fetch_imagesr8   )rT   r%   rP   s      r"   _prepare_images_structurez1SuperGlueImageProcessor._prepare_images_structure   s     ""6*.v66r$   torch.Tensor	do_resizesizeresamplez7PILImageResampling | tvF.InterpolationMode | int | None
do_rescalerescale_factordisable_groupingreturn_tensorsrD   c
                 (   t        ||      \  }}i }|j                         D ]   \  }}|r| j                  |||      }|||<   " t        ||      }t        ||      \  }}i }|j                         D ]+  \  }}|r| j	                  ||      }|	rt        |      }|||<   - t        ||      }t        dt        |      d      D cg c]
  }|||dz     }}|D cg c]  }t        j                  |d       }}t        d|i|      S c c}w c c}w )N)rb   )r^   r_   r   r(   )dimpixel_values)datatensor_type)r   itemsresizer   rescalerA   ranger   r;   stackr   )rT   r%   r]   r^   r_   r`   ra   rb   rc   rD   rP   grouped_imagesgrouped_images_indexprocessed_images_groupedr   stacked_imagesresized_imagesprocessed_imagesiimage_pairspairstacked_pairss                         r"   _preprocessz#SuperGlueImageProcessor._preprocess   sS    0EV^n/o,,#% %3%9%9%; 	=!E>!%^$QY!Z.<$U+	= ((@BVW/D^fv/w,,#% %3%9%9%; 	=!E>!%nn!M!5n!E.<$U+	= **BDXY =B!SIYEZ\]<^_q'AE2__ ?JJdTq1JJ .-!@n]] ` Ks   D
Doutputsr   target_sizes	thresholdc                    |j                   j                  d   t        |      k7  rt        d      t	        d |D              st        d      t        |t              r,t        j                  ||j                   j                        }n1|j                  d   dk7  s|j                  d   dk7  rt        d      |}|j                  j                         }||j                  d      j                  dddd      z  }|j                  t        j                        }g }t!        |j                   ||j"                  d	d	df   |j$                  d	d	df         D ]v  \  }}}	}
|d   dkD  }|d   dkD  }|d   |   }|d   |   }|	|   }|
|   }||kD  |dkD  z  ||j                  d   k  z  }||   }|||      }||   }|j'                  |||d
       x |S )a  
        Converts the raw output of [`SuperGlueKeypointMatchingOutput`] into lists of keypoints, scores and descriptors
        with coordinates absolute to the original image sizes.
        Args:
            outputs ([`SuperGlueKeypointMatchingOutput`]):
                Raw outputs of the model.
            target_sizes (`torch.Tensor` or `list[tuple[tuple[int, int]]]`, *optional*):
                Tensor of shape `(batch_size, 2, 2)` or list of tuples of tuples (`tuple[int, int]`) containing the
                target size `(height, width)` of each image in the batch. This must be the original image size (before
                any processing).
            threshold (`float`, *optional*, defaults to `0.0`):
                Threshold to filter out the matches with low scores.
        Returns:
            `list[Dict]`: A list of dictionaries, each dictionary containing the keypoints in the first and second image
            of the pair, the matching scores and the matching indices.
        r   zRMake sure that you pass in as many target sizes as the batch dimension of the maskc              3   8   K   | ]  }t        |      d k(    yw)r(   N)r   )r-   target_sizes     r"   r.   zISuperGlueImageProcessor.post_process_keypoint_matching.<locals>.<genexpr>   s     I[3{#q(Is   zTEach element of target_sizes must contain the size (h, w) of each image of the batch)devicer   r(   N)
keypoints0
keypoints1matching_scores)maskr   r   r6   r4   r2   r3   r;   tensorr   	keypointscloneflipreshapetoint32zipmatchesr   append)rT   ry   rz   r{   image_pair_sizesr   results	mask_pairkeypoints_pairr   scoresmask0mask1r   r   matches0scores0valid_matchesmatched_keypoints0matched_keypoints1r   s                        r"   post_process_keypoint_matchingz6SuperGlueImageProcessor.post_process_keypoint_matching   s   , <<a C$55qrrILIIsttlD)$||LATATU!!!$)\-?-?-Ba-G j   ,%%++-	 0 5 5b 9 A A"aA NN	LL-	:=LL)W__QT%:G<S<STUWXTX<Y;
 	6I~w aL1$EaL1$E'*51J'*51Ju~HUmG %y0X]CxR\RbRbcdReGefM!+M!:!+H],C!D%m4ONN"4"4'6#	2 r$   keypoint_matching_outputzImage.Imagec           	      <   t        |      }|D cg c]  }t        |       }}t        dt        |      d      D cg c]
  }|||dz     }}g }t	        ||      D ]  \  }}|d   j
                  dd \  }	}
|d   j
                  dd \  }}t        j                  t        |	|      |
|z   dft        j                        }t        j                  |d         |d|	d|
f<   t        j                  |d         |d||
df<   t        j                  |j                               }t        j                  |      }|d   j!                  d      \  }}|d   j!                  d      \  }}t	        |||||d	         D ]  \  }}}}}| j#                  |      }|j%                  ||||
z   |f|d
       |j'                  |dz
  |dz
  |dz   |dz   fd       |j'                  ||
z   dz
  |dz
  ||
z   dz   |dz   fd        |j)                  |        |S c c}w c c}w )a  
        Plots the image pairs side by side with the detected keypoints as well as the matching between them.

        Args:
            images:
                Image pairs to plot. Same as `EfficientLoFTRImageProcessor.preprocess`. Expects either a list of 2
                images or a list of list of 2 images list with pixel values ranging from 0 to 255.
            keypoint_matching_output (List[Dict[str, torch.Tensor]]]):
                A post processed keypoint matching output

        Returns:
            `List[PIL.Image.Image]`: A list of PIL images, each containing the image pairs side by side with the detected
            keypoints as well as the matching between them.
        r   r(   Nr   r   )dtyper   r   r   )fillrO   black)r   )r8   r   rl   r   r   r   r;   zerosmaxuint8
from_numpyr   	fromarraynumpyr   Drawunbind
_get_colorlineellipser   )rT   r%   r   r!   rt   ru   r   r5   pair_outputheight0width0height1width1
plot_imageplot_image_pildrawkeypoints0_xkeypoints0_ykeypoints1_xkeypoints1_ykeypoint0_xkeypoint0_ykeypoint1_xkeypoint1_ymatching_scorecolors                             r"   visualize_keypoint_matchingz3SuperGlueImageProcessor.visualize_keypoint_matching  sh   ( 185;<E.'<<273v;2JKQva!a%(KK'*;8P'Q 	+#J(m11"15OGV(m11"15OGVc'7&;Vf_a%PX]XcXcdJ,1,<,<Z],KJxx&(),1,<,<Z],KJxx()"__Z-=-=-?@N>>.1D)4\)B)I)I!)L&L,)4\)B)I)I!)L&L,VYlL,TeHfW R[+{N 7		 +{V/C[Q  
 kAo{QaQ\_`Q`ahop 6)A-{Qf@TWX@XZehiZij    NN>*7	+8 A =Ks
   HHc                 N    t        dd|z
  z        }t        d|z        }d}|||fS )zMaps a score to a color.   r   r   )int)rT   scorergbs        r"   r   z"SuperGlueImageProcessor._get_color=  s3    q5y!"e!Qwr$   )T)g        )%rE   rF   rG   rC   valid_kwargsr   BILINEARr_   r^   default_to_squarer]   r`   ra   do_normalizerD   r   rS   r   r	   r   rX   r[   r3   rI   r   floatstrr   rx   tupledictr;   Tensorr   r   r   __classcell__)rU   s   @r"   rM   rM   {   s   0L!**HC(DIJNLL#(E!F # 4 4v>[7\ 4am 4 477 
	7& ")^^$)^ )^ 	)^
 L)^ )^ )^ +)^ j(4/)^ )^ 
)^^ 	B2B !4;.B 	B
 
d3$%	&BH5 #'tC,='>"?5 
m		5nr$   rM   )r!   r\   )r!   r\   rV   r\   )*rH   typingr   r;   image_processing_backendsr   image_processing_utilsr   image_transformsr   r   image_utilsr	   r
   r   r   r   r   r   r   processing_utilsr   r   utilsr   r   r   modeling_supergluer   r   r   r   torchvision.transforms.v2r   r?   r#   r8   r<   rA   rC   rM   __all__r+   r$   r"   <module>r      s    +    ; 2 E	 	 	 5  C$ 7$J $8>>>&L  F0 F FR %
%r$   