
    ^Nj1              
       6   U d dl mZmZmZmZ d dlZd dlmZ d dl	m
Z
mZ d dlmZ d dlmZmZ d dlmZmZ d dlmZ d d	lmZmZmZ d d
lmZmZ  eddddd ed      dgd      gZee   ed<    G d deee         Z  G d dee         Z! G d dee         Z"y)    )AnyIterableSequenceTypeN)Encoding)OnnxProvider
ImageInput)OnnxOutputContext)
NumpyArrayDevice)define_cache_dir
iter_batch)&LateInteractionMultimodalEmbeddingBase)OnnxMultimodalModelTextEmbeddingWorkerImageEmbeddingWorker)DenseModelDescriptionModelSourcezQdrant/colpali-v1.3-fp16   z[Text embeddings, Multimodal (text&image), English, 50 tokens query length truncation, 2024.mitg      @)hfzmodel.onnx_dataz
model.onnx)modeldimdescriptionlicense
size_in_GBsourcesadditional_files
model_filesupported_colpali_modelsc                       e Zd ZdZdZdZddgZdZ ej                  dgdz  g d	z         Z
 ej                  d
gdz        Zdddej                  ddddfdededz  dedz  dee   dz  deez  dee   dz  dededz  dedz  def fdZedee   fd       Zd.dZdedee   fdZdedee   fdZdee   dedee   fd Z 	 	 d/d!eee   z  d"ed#ededef
d$Z!d%e"eef   dede"eef   fd&Z#d%e"eejH                  f   dede"eef   fd'Z%	 	 d0deee   z  d"ed(edz  dedee   f
d)Z&	 	 d1d*e'ee'   z  d"ed(edz  dedee   f
d+Z(ede)e*e      fd,       Z+ede)e,e      fd-       Z- xZ.S )2ColPalizQuery: z<s>z<pad>   i  )     r%   i    )r#   i!  i=  ip	  i l      i  NF
model_name	cache_dirthreads	providerscuda
device_ids	lazy_load	device_idspecific_model_pathkwargsc
                 6   t        |   |||fi |
 || _        || _        | j	                  |
      | _        || _        || _        d| _        ||| _        n | j                  | j                  d   | _        | j                  |      | _
        t        t        |            | _        |	| _        | j                  | j                  | j                  | j                   | j                        | _        d| _        d| _        | j                  s| j)                          yy)a"  
        Args:
            model_name (str): The name of the model to use.
            cache_dir (str, optional): The path to the cache directory.
                                       Can be set using the `FASTEMBED_CACHE_PATH` env variable.
                                       Defaults to `fastembed_cache` in the system's temp directory.
            threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
            providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
                Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
            cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`
                Defaults to Device.AUTO.
            device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
                workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive
                with `providers`. Defaults to None.
            lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
                Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
            device_id (Optional[int], optional): The device id to use for loading the model in the worker process.

        Raises:
            ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
        Nr   )local_files_onlyr1   )super__init__r,   r/   _select_exposed_session_options_extra_session_optionsr.   r-   r0   _get_model_descriptionmodel_descriptionstrr   r*   _specific_model_pathdownload_model_local_files_only
_model_dirmask_token_idpad_token_idload_onnx_model)selfr)   r*   r+   r,   r-   r.   r/   r0   r1   r2   	__class__s              ~/Users/ahmed/devFolder/Ultron/claude-voice/.venv/lib/python3.12/site-packages/fastembed/late_interaction_multimodal/colpali.pyr6   zColPali.__init__.   s   F 	YB6B""&*&J&J6&R# %	 &* &DN__(!__Q/DN!%!<!<Z!H-i89$7!--""NN!33 $ 9 9	 . 
 " ~~  "     returnc                     t         S )zLists the supported models.

        Returns:
            list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
        )r    clss    rE   _list_supported_modelszColPali._list_supported_modelsq   s
     ('rF   c           	          | j                  | j                  | j                  j                  | j                  | j
                  | j                  | j                  | j                         y )N)	model_dirr   r+   r,   r-   r0   extra_session_options)	_load_onnx_modelr?   r:   r   r+   r,   r-   r0   r8   )rC   s    rE   rB   zColPali.load_onnx_modelz   sP    oo--88LLnnnn"&"="= 	 	
rF   outputc                     | j                   j                  J d       |j                  j                  |j                  j                  d   d| j                   j                        S )  
        Post-process the ONNX model output to convert it into a usable format.

        Args:
            output (OnnxOutputContext): The raw output from the ONNX model.

        Returns:
            Iterable[NumpyArray]: Post-processed output as NumPy arrays.
        zModel dim is not definedr   )r:   r   model_outputreshapeshaperC   rP   s     rE   _post_process_onnx_image_outputz'ColPali._post_process_onnx_image_output   s_     %%))5Q7QQ5""**%%a("d.D.D.H.H
 	
rF   c                     |j                   S )rR   )rT   rW   s     rE   _post_process_onnx_text_outputz&ColPali._post_process_onnx_text_output   s     """rF   	documentsc                     g }|D ]D  }| j                   | j                  z   |z   | j                  dz  z   }|dz  }|j                  |       F | j                  j                  |      }|S )N
   
)	BOS_TOKENQUERY_PREFIX	PAD_TOKENappend	tokenizerencode_batch)rC   r[   r2   texts_queryqueryencodeds         rE   tokenizezColPali.tokenize   sk    !#ENNT%6%66>RTATTETMEu%	 
 ..--k:rF   texts
batch_sizeinclude_extensionc           
      ~   t        | d      r| j                  | j                          d}t        |t              r|gn|}| j
                  J |r| j                  n| j
                  j                  }t        ||      D ]7  }|t         ||      D cg c]  }t        |j                         c}      z  }9 |S c c}w )Nr   r   )hasattrr   rB   
isinstancer;   rc   rh   rd   r   sumattention_mask)	rC   ri   rj   rk   r2   	token_numtokenize_funcbatchencodings	            rE   token_countzColPali.token_count   s     tW%);  "	%eS1u~~)))):@[@[z2E=Y^K_`K_xc("9"9:K_`aaI 3 as   B:
onnx_inputc           	      X   t        j                  |d   D cg c]"  }| j                  |dd  j                         z   $ c}      |d<   t        j                  | j
                  t         j                        }t        j                  |d   D cg c]  }| c}      |d<   |S c c}w c c}w )N	input_idsr#   )dtypepixel_values)nparrayQUERY_MARKER_TOKEN_IDtolistzerosIMAGE_PLACEHOLDER_SIZEfloat32)rC   rv   r2   rx   empty_image_placeholder_s         rE   _preprocess_onnx_text_inputz#ColPali._preprocess_onnx_text_input   s     #%(( ",K!8!8I **Yqr]-A-A-CC!8#

; /1hh''rzz/
 &(XX.8.EF.E$.EF&

>"  Gs   'B"	B'c                     t        j                  |d   D cg c]  }| j                   c}      |d<   t        j                  |d   D cg c]  }| j                   c}      |d<   |S c c}w c c}w )a2  
        Add placeholders for text input when processing image data for ONNX.
        Args:
            onnx_input (Dict[str, NumpyArray]): Preprocessed image inputs.
            **kwargs: Additional arguments.
        Returns:
            Dict[str, NumpyArray]: ONNX input with text placeholders.
        rz   rx   rp   )r{   r|   EMPTY_TEXT_PLACEHOLDEREVEN_ATTENTION_MASK)rC   rv   r2   r   s       rE   _preprocess_onnx_image_inputz$ColPali._preprocess_onnx_image_input   s     #%((2<^2LM2LQT((2LM#

; (*xx/9./IJ/I!T%%/IJ(

#$  N Ks   A)
A.parallelc              +     K    | j                   d| j                  t        | j                        |||| j                  | j
                  | j                  | j                  | j                  | j                  d|E d{    y7 w)ac  
        Encode a list of documents into list of embeddings.

        Args:
            documents: Iterator of documents or single document to embed
            batch_size: Batch size for encoding -- higher values will use more memory, but be faster
            parallel:
                If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
                If 0, use all available cores.
                If None, don't use data-parallel processing, use default onnxruntime threading instead.

        Returns:
            List of embeddings, one per document
        )r)   r*   r[   rj   r   r,   r-   r.   r4   r1   rN   N )
_embed_documentsr)   r;   r*   r,   r-   r.   r>   r<   r8   )rC   r[   rj   r   r2   s        rE   
embed_textzColPali.embed_text   s}     * )4(( 
$..)!nn!33 $ 9 9"&"="=
 
 	
 	
   A;B=B>Bimagesc              +     K    | j                   d| j                  t        | j                        |||| j                  | j
                  | j                  | j                  | j                  | j                  d|E d{    y7 w)aa  
        Encode a list of images into list of embeddings.

        Args:
            images: Iterator of image paths or single image path to embed
            batch_size: Batch size for encoding -- higher values will use more memory, but be faster
            parallel:
                If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
                If 0, use all available cores.
                If None, don't use data-parallel processing, use default onnxruntime threading instead.

        Returns:
            List of embeddings, one per document
        )r)   r*   r   rj   r   r,   r-   r.   r4   r1   rN   Nr   )
_embed_imagesr)   r;   r*   r,   r-   r.   r>   r<   r8   )rC   r   rj   r   r2   s        rE   embed_imagezColPali.embed_image	  s}     * &4%% 
$..)!nn!33 $ 9 9"&"="=
 
 	
 	
r   c                     t         S N)ColPaliTextEmbeddingWorkerrI   s    rE   _get_text_worker_classzColPali._get_text_worker_class-  s    ))rF   c                     t         S r   )ColPaliImageEmbeddingWorkerrI   s    rE   _get_image_worker_classzColPali._get_image_worker_class1  s    **rF   )rG   N)r&   F)   N)   N)/__name__
__module____qualname__r`   r_   ra   r}   r   r{   r|   r   r   r   AUTOr;   intr   r   boollistr   r6   classmethodr   rK   rB   r
   r   r   rX   rZ   r   rh   ru   dictr   ndarrayr   r   r	   r   r   r   r   r   r   __classcell__)rD   s   @rE   r"   r"   "   s(   LIII*%RXX	4<< #"((A3:.
 !%"37$kk'+ $*.A#A# :A# t	A#
 L)D0A# VmA# I$A# A# :A# !4ZA# A#F (t,A'B ( (	

!
 
*	
$#!# 
*	#$s) s tH~  "'	Xc]"   	
  
"sJ/;>	c:o	"sBJJ/;>	c:o	, #	"
#&"
 "
 *	"

 "
 
*	"
N #	"
Xj11"
 "
 *	"

 "
 
*	"
H *t,?
,K'L * * +-A*-M(N + +rF   r"   c                   $    e Zd ZdedededefdZy)r   r)   r*   r2   rG   c                      t        d||dd|S Nr(   )r)   r*   r+   r   r"   rC   r)   r*   r2   s       rE   init_embeddingz)ColPaliTextEmbeddingWorker.init_embedding7  '     
!
 	
 	
rF   Nr   r   r   r;   r   r"   r   r   rF   rE   r   r   6  $    
 
 
 
PW 
rF   r   c                   $    e Zd ZdedededefdZy)r   r)   r*   r2   rG   c                      t        d||dd|S r   r   r   s       rE   r   z*ColPaliImageEmbeddingWorker.init_embeddingA  r   rF   Nr   r   rF   rE   r   r   @  r   rF   r   )#typingr   r   r   r   numpyr{   
tokenizersr   fastembed.commonr   r	   fastembed.common.onnx_modelr
   fastembed.common.typesr   r   fastembed.common.utilsr   r   Pfastembed.late_interaction_multimodal.late_interaction_multimodal_embedding_baser   ;fastembed.late_interaction_multimodal.onnx_multimodal_modelr   r   r   "fastembed.common.model_descriptionr   r   r    r   __annotations__r"   r   r   r   rF   rE   <module>r      s    0 0   5 9 5 ? 
 R (q9:+,	9 $45 Q+46I*6U Q+h
!4Z!@ 

"6z"B 
rF   