
    ^js`                        d dl Z d dlmZ d dlmZmZmZ d dlZd dl	Z	ddl
mZ ddlmZmZ ddlmZmZmZmZ ddlmZmZmZ dd	lmZ dd
lmZ erddlmZ  G d ded      Z ed      e G d de                    ZdgZy)    N)
accumulate)TYPE_CHECKINGOptionalUnion   )BatchFeature)
ImageInputis_valid_image)MultiModalDataProcessingKwargsProcessorMixinUnpack)
AddedTokenBatchEncoding	TextInput)auto_docstring)requires)PreTokenizedInputc                   (    e Zd ZddiddddddidZy	)
ColModernVBertProcessorKwargspaddinglongestTchannels_first)return_row_col_infodata_formatdo_convert_rgbreturn_tensorspt)text_kwargsimages_kwargscommon_kwargsN)__name__
__module____qualname__	_defaults     /var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/transformers/models/colmodernvbert/processing_colmodernvbert.pyr   r   (   s/     y
 $(+"

 +D1
Ir'   r   F)total)torch)backendsc                   "    e Zd ZeZ	 	 	 	 	 d dededz  dedz  f fdZe	 	 	 d!de	e
e	   z  e
e
e	      z  deede
e   e
d   f   dedz  d	ee   d
ef
d       Z	 	 d"de	dz  deede
e   e
d   f   d	ee   fdZ	 	 d"de	dz  deede
e   e
d   f   d	ee   f fdZdeded
efdZde
de
e   d
e
e
e      fdZd#dZ	 d#de	dz  d	ee   d
efdZdee
e   z  d	ee   d
efdZ	 	 	 d$dede
d   f   dede
d   f   deded   dedef   d
dfdZ xZS )%ColModernVBertProcessorNimage_seq_lenvisual_prompt_prefixquery_prefixc                    d}t        ddd      j                  | _        t        ddd      j                  | _        t        ddd      j                  | _        d| _        || _        |j                  | j                        | _        |j                  | j                        | _	        |j                  | j
                        | _
        t        d	      D 	cg c]0  }t        d	      D ]   }	|j                  d
|dz    d|	dz    d      " 2 c}	}| _        t        j                  d      | _        d| j                  | j                  | j                  gi}
|j!                  |
       |j                  | j                        | _        t#        | H  ||fd|i| |xs d| j                   d| _        |xs d| _        | j                  | _        yc c}	}w )a  
        image_seq_len (`int`, *optional*, defaults to 64):
            The length of the image sequence i.e. the number of <image> tokens per image in the input.
        visual_prompt_prefix (`str`, *optional*):
            A string that gets tokenized and prepended to the image tokens.
        query_prefix (`str`, *optional*):
            A prefix to be used for the query.
        Nz<fake_token_around_image>FT)
normalizedspecialz<image>z<end_of_utterance>z<global-img>   <row_   _col_>z*(\n?<global-img>\n?|<row_\d+_col_\d+>\n?)+additional_special_tokenschat_templatez<|begin_of_text|>User:z0Describe the image.<end_of_utterance>
Assistant: )r   contentfake_image_tokenimage_tokenend_of_utterance_tokenglobal_image_tagr.   convert_tokens_to_idsimage_token_idfake_image_token_idglobal_image_token_idrangerow_col_idsrecompile%_regex_to_remove_extra_special_tokensadd_special_tokenssuper__init__r/   r0   query_augmentation_token)selfimage_processor	tokenizerr:   r.   r/   r0   kwargsijtokens_to_add	__class__s              r(   rL   z ColModernVBertProcessor.__init__;   s   $  *+FSXbf g o o%iE4PXX&01ERWae&f&n&n# .*'==d>N>NO#,#B#B4CXCX#Y %.%D%DTEZEZ%["SXYZS[
NOejklem
`aI++eAE7%Awa,HI
I
 68ZZ@m5n2 (%%  ++*
 	$$]3'==d>N>NO)[=[TZ[$8 %
$T%5%5$66gh 	! ).B(,(C(C%1
s   5Gimagestextr   rQ   returnc                     | j                   d||d|\  }} | j                  d||d|  | j                  t        fd| j                  j
                  i|}||n| j                  }|d   j                  dd      }|d   j                  dd      }|d   j                  dd      }i x}	}
| | j                  |fi |d	   \  }	}|	j                  d
d       |	j                  dd       || j                  ||      \  }} | j                  |fi |d   }
|r||
d<   g }t        |      D ]e  \  }}g }|D ]H  }|d   \  }}|
j                  ||      }|
j                  ||dz
        }|j                  ||z
  dz          J |j                  |       g |r| j                  |
d   |      |
d<   | j                  ||
dg       n| | j                  dd|i|d   }
t        i |
|	|      S )a  
        image_seq_len (`int`, *optional*):
            The length of the image sequence. If not provided, the default value of self.image_seq_len is used.
            image_seq_len should be equal to int(((image_size // patch_size) ** 2) / (scale_factor**2))
        )rV   rW   tokenizer_init_kwargsNr   return_text_replacement_offsetsFreturn_mm_token_type_idsr   r    rowscols)images_replacementstext_replacement_offsetsnew_spanr6   	input_idsmm_token_type_idsimage)
modalitiesrW   )datatensor_typer&   )prepare_inputs_layoutvalidate_inputs_merge_kwargsr   rP   init_kwargsr.   pop_process_imagesget_text_with_replacements	enumeratechar_to_tokenappendcreate_mm_token_type_ids_check_special_mm_tokensr   )rN   rV   rW   r.   rQ   output_kwargsr[   r\   r   image_inputstext_inputsr_   r`   batch_image_seq_lengthsbatch_idtext_replacement_offsetimage_seq_lensrf   startendstart_id_pos
end_id_poss                         r(   __call__z ColModernVBertProcessor.__call__p   s    2t11UdUfU@F@@***)
"&.."<"<
 
 *7)BHZHZ*7*F*J*JKlns*t'#0#?#C#CD^`e#f &}599:JDQ%''{0D0D0DV0n}]lOm0n-L- VT*VT*151P1P.A 2Q 2.. -dnnTR]=5QR2>VK :;*,'9BC[9\ C5H5%'N 7 M%)*%5
s'2'@'@5'Q%0%>%>xq%Q
&--j<.G!.KLM ,22>BC ,7;7T7T#K02I8K 34 --dKWI-V($..SdSmM6RSK!@K!@<!@n]]r'   c                 B   |#t        |t              r|g}|j                         }|| j                  j	                  |      }t        |      r|gg}||fS t        |t        t        f      rt        |d         r||D cg c]  }|j                  | j                         }}dgt        t        |            z   }t        t        |            D cg c]  }|||   ||dz        }}t        |      |d   kD  r|||d   d  gz   }||fS |}||fS |g}||fS c c}w c c}w )Nr   r6   )
isinstancestrcopyrO   fetch_imagesr
   listtuplecountr>   r   rE   len)	rN   rV   rW   rQ   samplen_images_in_textcumsum_images_in_textrR   split_imagess	            r(   rh   z-ColModernVBertProcessor.prepare_inputs_layout   sk    $$v99;D))66v>Ff%!($ t|# FT5M2~fQi7P#UY'Z6T5E5E(F'Z$'Z-.C$zBR7S2T,T) "'s+;'<!=$ 4Q7:OPQTUPU:VW$L $
 6{%:2%>>!-8Mb8Q8S1T0U!U t|	 ". t| %XFt| ([$s    "DDc                    t        |   ||fi | ||t        d      ||D cg c]  }|j                  | j                         }}|J|D cg c]  }t        |       }}||k7  r,t        d| j                   d| d| j                   d| d	      y |1t        |      r%t        dt        |       d| j                   d      y y y c c}w c c}w )	Nz+You must provide either `text` or `images`.zThe total number of zP tokens in the prompts should be the same as the number of images passed. Found  z tokens and z images per sample.zFound z. tokens in the text but no images were passed.)rK   ri   
ValueErrorr   r>   r   anysum)	rN   rV   rW   rQ   r   r   sublistn_images_in_imagesrU   s	           r(   ri   z'ColModernVBertProcessor.validate_inputs   s    	77<FNJKKMQR6T-=-= >RR!BH%Iwc'l%I"%I#'99$.t/?/?.@ A""2!31T5E5E4FlSeRffy{  :
 C(8$9 S!1231T5E5E4FFtu  %: R%Is   "CCru   	image_idxc           	         |d   D cg c]  }|D ]  }|  c}}|   }|d   D cg c]  }|D ]  }|  c}}|   }|dk(  rI|dk(  rD| j                    | j                   z   | j                   | j                  z  z   | j                    z   S d}	t	        |      D ]R  }
t	        |      D ]=  }|	| j                    d|
dz    d|dz    dz   | j                   | j                  z  z   z  }	? |	d	z  }	T |	d	| j                    | j                   z   | j                   | j                  z  z   | j                    z   z  }	|	S c c}}w c c}}w )
Nr]   r^   r   r;   r5   r6   r7   r8   
)r=   r@   r>   r.   rE   )rN   ru   r   row_listrow
image_rowscol_listcol
image_colstext_split_imagesn_hn_ws               r(   replace_image_tokenz+ColModernVBertProcessor.replace_image_token   s   *6v*>Sh(S3cScST]^
*6v*>Sh(S3cScST]^
?zQ(()**+-%%&$*<*<<= **+- !#Z( * , C%001!#'%ay:;!--.$2D2DDE% "T)!* T**+,**+-%%&$*<*<<= **+- %$5 TSs
   D:E rb   rw   c                    g }t        |      D ]  \  }}t        j                  ||         }t        j                  |      }t        j                  || j
                  k(        d   }d}	|D ]7  }
|	t        |      k\  r n'||	   }||
z   }d||| t        j                  ||      }	9 |j                  |j                                 |S )Nr   r6   )
ro   nparray
zeros_likewhererC   r   searchsortedrq   tolist)rN   rb   rw   rc   rR   seq_lengths	array_idsmm_token_typesimage_start_positionsrS   seq_lenr{   r|   s                r(   rr   z0ColModernVBertProcessor.create_mm_token_type_ids	  s     '(?@ 	>NA{1.I]]95N$&HHY$:R:R-R$STU$V!A& @122-a0go,-uS)OO$93?@ $$^%:%:%<=	> ! r'   c                     i }|t         j                  j                  di       }|j                  |       |D cg c]   } | j                  j
                  g || " }}| j                  dz   }| j                  dz   }t        | j                  dd      d         dz
  }	g }
g }|D ]B  \  }}}||z  dz   }|d	kD  r|	nd	}|
j                  ||z   ||z  z          |j                  |       D |j                  |
|d
       t        di |S c c}w )a  
        Computes the number of placeholder tokens needed for multimodal inputs with the given sizes.

        Args:
            image_sizes (`list[list[int]]`, *optional*):
                The input sizes formatted as (height, width) per each image.

        Returns:
            `MultiModalData`: A `MultiModalData` object holding number of tokens per each of the provided
            input modalities, along with other useful data.
        r    r      z

F)rJ   rb   r6   r   )num_image_tokensnum_image_patchesr&   )r   r%   getupdaterO   get_number_of_image_patchesr.   r   rP   rq   r   )rN   image_sizesrQ   vision_datar    
image_sizenum_image_row_colsbase_image_length
col_lengthextra_split_newliner   r   num_patchesnum_rowsnum_cols
row_lengthsplit_extras                    r(   _get_num_multimodal_tokensz2ColModernVBertProcessor._get_num_multimodal_tokens  sR    "9CCGGY[\M  ( #." A$$@@\*\m\" "
 !% 2 2Q 6++a/J
 #&dnnVPUn&VWb&c"dgh"h! "3E 6/Xx'(2Q6
5=\1q ''(9K(G:X`K`(ab!((5	6 4D[lmn,,,/"s   %Dc                 v    | j                   t        fd| j                  j                  i|}|d   j	                  dd      }|du}t        |      r|g}n^t        |t              rt        |d         rn?t        |t              r$t        |d   t              rt        |d   d         st        d      |D cg c]  }|j                  d       }}| j                  | j                  gt        |      z  ||d   |d   	      }|r.|d
   j                  |d   dk(  d      }|j                  d|i       |S c c}w )a  
        Prepare for the model one or several image(s). Handles input validation, RGB conversion,
        and prepends the `visual_prompt_prefix` to each image. Optionally computes labels from
        `token_type_ids` when a `suffix` is provided in `text_kwargs`.

        Args:
            images (`PIL.Image.Image`, `np.ndarray`, `torch.Tensor`, `list[PIL.Image.Image]`, `list[np.ndarray]`, `list[torch.Tensor]`):
                The image or batch of images to be prepared. Each image can be a PIL image, NumPy array or PyTorch
                tensor. In case of a NumPy array/PyTorch tensor, each image should be of shape (C, H, W), where C is a
                number of channels, H and W are image height and width.
            return_tensors (`str` or [`~utils.TensorType`], *optional*):
                If set, will return tensors of a particular framework. Acceptable values are:

                - `'pt'`: Return PyTorch `torch.Tensor` objects.
                - `'np'`: Return NumPy `np.ndarray` objects.

        Returns:
            [`BatchFeature`]: A [`BatchFeature`] with the following fields:

            - **input_ids** -- List of token ids to be fed to a model.
            - **attention_mask** -- List of indices specifying which tokens should be attended to by the model (when
              `return_attention_mask=True` or if *"attention_mask"* is in `self.model_input_names` and if `text` is not
              `None`).
            - **pixel_values** -- Pixel values to be fed to a model. Returned when `images` is not `None`.
        rZ   r   suffixNr   zAimages must be an image, list of images or list of list of imagesRGBr    )rW   rV   r    r   rb   token_type_idsilabels)rj   r   rP   rk   rl   r
   r   r   r   convertr   r/   r   masked_fillr   )	rN   rV   rQ   rt   r   return_token_type_idsrd   	batch_docr   s	            r(   process_imagesz&ColModernVBertProcessor.process_imagesI  s[   < +**)
"&.."<"<
 
 }-11(DA &d 2 &!XF%.*CVT*z&)T/J~^def^ghi^jOk`aa 5;;5%--&;; MM++,s6{:'8%m4	 " 
	 !{+77	BR8SWX8XZ^_Fh/0 <s   8D6c                     | j                   t        fd| j                  j                  i|}|d   j	                  dd      }t        |t              r|g}n.t        |t              rt        |d   t              st        d      || j                  dz  }|D cg c]  }| j                  |z   |z    }}| j                  |d|d   	      }|S c c}w )
ad  
        Prepare for the model one or several text queries. Handles input validation, prepends the
        `query_prefix`, and appends query augmentation tokens (used to pad query embeddings for
        better late-interaction retrieval performance).

        Args:
            text (`str`, `list[str]`, `list[list[str]]`):
                The sequence or batch of sequences to be encoded. Each sequence can be a string or a list of strings
                (pretokenized string). If the sequences are provided as list of strings (pretokenized), you must set
                `is_split_into_words=True` (to lift the ambiguity with a batch of sequences).
            return_tensors (`str` or [`~utils.TensorType`], *optional*):
                If set, will return tensors of a particular framework. Acceptable values are:

                - `'pt'`: Return PyTorch `torch.Tensor` objects.
                - `'np'`: Return NumPy `np.ndarray` objects.

        Returns:
            [`BatchFeature`]: A [`BatchFeature`] with the following fields:

            - **input_ids** -- List of token ids to be fed to a model.
            - **attention_mask** -- List of indices specifying which tokens should be attended to by the model (when
              `return_attention_mask=True` or if *"attention_mask"* is in `self.model_input_names` and if `text` is not
              `None`).
        rZ   r   r   Nr   z*Text must be a string or a list of strings
   F)rW   r   r   )rj   r   rP   rk   rl   r   r   r   r   rM   r0   r   )rN   rW   rQ   rt   r   querytexts_querybatch_querys           r(   process_queriesz'ColModernVBertProcessor.process_queries  s    : +**)
"&.."<"<
 
 }-11(DAdC 6DT4(ZQ-EIJJ >22R7F SW!W$"3"3e";f"D!W!Wmm"'%m4 $ 
  "Xs   Cquery_embeddingsztorch.Tensorpassage_embeddings
batch_sizeoutput_dtypeztorch.dtypeoutput_deviceztorch.devicec           	         t        |      dk(  rt        d      t        |      dk(  rt        d      |d   j                  |d   j                  k7  rt        d      |d   j                  |d   j                  k7  rt        d      ||d   j                  }g }t	        dt        |      |      D ]%  }g }t
        j                  j                  j                  j                  ||||z    dd      }	t	        dt        |      |      D ]  }
t
        j                  j                  j                  j                  ||
|
|z    dd      }|j                  t        j                  d|	|      j                  d	
      d   j                  d
              |j                  t        j                  |d
      j                  |      j                  |             ( t        j                  |d
      S )a[  
        Compute the late-interaction/MaxSim score (ColBERT-like) for the given multi-vector
        query embeddings (`qs`) and passage embeddings (`ps`). For ColQwen2, a passage is the
        image of a document page.

        Because the embedding tensors are multi-vector and can thus have different shapes, they
        should be fed as:
        (1) a list of tensors, where the i-th tensor is of shape (sequence_length_i, embedding_dim)
        (2) a single tensor of shape (n_passages, max_sequence_length, embedding_dim) -> usually
            obtained by padding the list of tensors.

        Args:
            query_embeddings (`Union[torch.Tensor, list[torch.Tensor]`): Query embeddings.
            passage_embeddings (`Union[torch.Tensor, list[torch.Tensor]`): Passage embeddings.
            batch_size (`int`, *optional*, defaults to 128): Batch size for computing scores.
            output_dtype (`torch.dtype`, *optional*, defaults to `torch.float32`): The dtype of the output tensor.
                If `None`, the dtype of the input embeddings is used.
            output_device (`torch.device` or `str`, *optional*, defaults to "cpu"): The device of the output tensor.

        Returns:
            `torch.Tensor`: A tensor of shape `(n_queries, n_passages)` containing the scores. The score
            tensor is saved on the "cpu" device.
        r   zNo queries providedzNo passages providedz/Queries and passages must be on the same devicez-Queries and passages must have the same dtypeT)batch_firstpadding_valuezbnd,csd->bcnsr   )dimr   r6   )r   r   devicedtyperE   r*   nnutilsrnnpad_sequencerq   einsummaxr   catto)rN   r   r   r   r   r   scoresrR   batch_scoresbatch_queriesrS   batch_passagess               r(   score_retrievalz'ColModernVBertProcessor.score_retrieval  s   @  A%233!"a'344A%%);A)>)E)EENOOA$$(:1(=(C(CCLMM+A.44L%'q#./< 	]A/1L!HHNN..;; Q^4$VW < M 1c"45zB !&!3!3!@!@&q1z>:\] "A " ##LL-PTTYZT[\]^bbghbi	 MM%))La8;;LILL][\	] yyQ''r'   )NN@   NN)NNN)NN)N)   Ncpu)r"   r#   r$   r   valid_processor_kwargsintr   rL   r   r	   r   r   r   r   r   r   rh   r   ri   dictr   rr   r   r   r   r   r   r   __classcell__)rU   s   @r(   r-   r-   6   s    ;
 +/#'3D
 3D "Dj3D Dj3Dj  JNbf$(	>^T*--T*5E0FF>^ I2DOTJ]E^^_>^ Tz	>^
 67>^ 
>^ >^D %)bf T!  I2DOTJ]E^^_  67	 H %)bfT! I2DOTJ]E^^_ )*	2% % % %:!$ !QUVYQZ !_cdhildm_n !*)-Z %)@T!@ 67@ 
	@D7$y/)7 677 
	7z 0449>(^0D DE>( ".$~2F"FG>( 	>(
 }->( ^S01>( 
>(r'   r-   ) rG   	itertoolsr   typingr   r   r   numpyr   r*   feature_extraction_utilsr   image_utilsr	   r
   processing_utilsr   r   r   r   tokenization_utils_baser   r   r   r   r   utils.import_utilsr   r   r   r-   __all__r&   r'   r(   <module>r      s~   * 
   1 1   4 5 X X K K # * <$4E  
:J(n J(  J(Z %
%r'   