
    ^Njh1                     P   U d dl Z d dlmZmZmZmZ d dl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 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  eddddd ed      d       eddddd ed      d      gZe e   e!d<    G d deee         Z" G d dee         Z#y)    N)AnyIterableSequenceType)Encoding	Tokenizer)load_tokenizer)
NumpyArrayDevice)OnnxProvider)OnnxOutputContext)define_cache_dir
iter_batch) LateInteractionTextEmbeddingBase)OnnxTextModelTextEmbeddingWorker)DenseModelDescriptionModelSourcezcolbert-ir/colbertv2.0   zQText embeddings, Unimodal (text), English, 512 input tokens truncation, 2023 yearmitg)\(?)hfz
model.onnx)modeldimdescriptionlicense
size_in_GBsources
model_filez%answerdotai/answerai-colbert-small-v1`   zQText embeddings, Unimodal (text), English, 512 input tokens truncation, 2024 yearz
apache-2.0gp=
ף?zvespa_colbert.onnxsupported_colbert_modelsc                   8    e Zd ZdZdZdZdZ	 d'dedede	de
e   fd	Z	 d'd
eeef   dede	deeef   fdZd'dee   dede	dee   fdZdedee   fdZdee   dee   fdZ	 	 	 d(dee
e   z  dededede	defdZedee   fd       Zdddej6                  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d)d"Z	 	 d*dee
e   z  ded#edz  de	de
e   f
d$Z dee
e   z  de	de
e   fd%Z!ede"e#e      fd&       Z$ xZ%S )+Colbert         z[MASK]outputis_dockwargsreturnc              +     K   |s|j                   D ]  }|  y |j                  |j                  t        d      t	        |j                        D ]G  \  }}t	        |      D ]4  \  }}|| j
                  v s|| j                  k(  s$d|j                  ||f<   6 I |xj                   t        j                  |j                  d      z  c_         t        j                  j                  |j                   ddd      }	t        j                  |	d      }
|xj                   |
z  c_         t        |j                   |j                        D ]  \  }}||dk(       y w)NzJinput_ids and attention_mask must be provided for document post-processingr   r$   T)ordaxiskeepdimsg-q=r#   )model_output	input_idsattention_mask
ValueError	enumerate	skip_listpad_token_idnpexpand_dimslinalgnormmaximumzip)selfr&   r'   r(   	embeddingitoken_sequencejtoken_idr8   norm_clampedr0   s               s/Users/ahmed/devFolder/Ultron/claude-voice/.venv/lib/python3.12/site-packages/fastembed/late_interaction/colbert.py_post_process_onnx_outputz!Colbert._post_process_onnx_output.   s>     #00	 1 '6+@+@+H `  &/v/?/?%@!>#,^#<KAx4>>1XARAR5R67--ad3 $= &A
 2>>&2G2G#KK99>>&"5"511t>TD::dE2L</-01D1DfF[F[-\)	>! 344 .]s   BECE
onnx_inputc                 *   |r| j                   n| j                  }t        j                  |d   j	                  t        j
                        d|d      |d<   t        j                  |d   j	                  t        j
                        ddd      |d<   |S )Nr/   r#   )r,   r0   )DOCUMENT_MARKER_TOKEN_IDQUERY_MARKER_TOKEN_IDr5   insertastypeint64)r;   rD   r'   r(   marker_tokens        rB   _preprocess_onnx_inputzColbert._preprocess_onnx_inputG   s     9?t44DD^D^"$)){#**2884aA#

; (*yy'(//91aa(

#$     	documentsc                 r    |r| j                  |      S | j                  t        t        |                  S )N)rN   )query)_tokenize_documents_tokenize_querynextiter)r;   rN   r'   r(   s       rB   tokenizezColbert.tokenizeS   s@      $$y$9	
 %%Di,A%B	
rM   rP   c                 Z    | j                   J | j                   j                  |g      }|S N)query_tokenizerencode_batch)r;   rP   encodeds      rB   rR   zColbert._tokenize_queryZ   s1    ##///&&33UG<rM   c                 <    | j                   j                  |      }|S rW   )	tokenizerrY   )r;   rN   rZ   s      rB   rQ   zColbert._tokenize_documents_   s    ..--i8rM   Ftexts
batch_sizeinclude_extensionc                    t        | d      r| j                  | j                          d}t        |t              r|gn|}|r| j
                  n| j                  }|J t        ||      D ]z  }|j                  |      D ]S  }	|r|t        |	j                        z  }t        |	j                        }
|r|t        |
| j                        z  }O||
z  }U |sm|t        |      z  }| |S )Nr   r   )hasattrr   load_onnx_model
isinstancestrr\   rX   r   rY   sumr0   maxMIN_QUERY_LENGTHlen)r;   r]   r^   r'   r_   r(   	token_numr\   batchtokensattend_counts              rB   token_countzColbert.token_countc   s     tW%);  "	%eS1u&,DNN$2F2F	$$$z2E#007V%:%:!;;I#&v'<'<#=L(!St7L7L%MM	 "\1	 8 !S 	 3  rM   c                     t         S )zLists the supported models.

        Returns:
            list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.
        )r    clss    rB   _list_supported_modelszColbert._list_supported_models   s
     ('rM   N
model_name	cache_dirthreads	providerscuda
device_ids	lazy_load	device_idspecific_model_pathc
                 b   t        |   |||fi |
 || _        || _        | j	                  |
      | _        || _        || _        d| _        ||| _        n | j                  | j                  d   | _        | j                  |      | _
        t        t        |            | _        |	| _        | j                  | j                  | j                  | j                   | j                        | _        d| _        d| _        t)               | _        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.
            specific_model_path (Optional[str], optional): The specific path to the onnx model dir if it should be imported from somewhere else

        Raises:
            ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
        Nr   )local_files_onlyrz   )super__init__ru   rx   _select_exposed_session_options_extra_session_optionsrw   rv   ry   _get_model_descriptionmodel_descriptionrd   r   rs   _specific_model_pathdownload_model_local_files_only
_model_dirmask_token_idr4   setr3   rX   rb   )r;   rr   rs   rt   ru   rv   rw   rx   ry   rz   r(   	__class__s              rB   r~   zColbert.__init__   s   H 	YB6B""&*&J&J6&R# %	 &* &DN__(!__Q/DN!%!<!<Z!H-i89$7!--""NN!33 $ 9 9	 . 
 *.(,#&515~~  " rM   c           	      j   | j                  | j                  | j                  j                  | j                  | j
                  | j                  | j                  | j                         t        | j                        \  | _
        }| j                  J | j                  | j                     | _        | j                  j                  d   | _        t"        j$                  D ch c],  }| j                  j'                  |d      j(                  d   . c}| _        | j                  j,                  d   }| j                  j/                  |dz
  	       | j                  j/                  |dz
  	       | j                  j1                  | j                  | j                  | j2                  
       y c c}w )N)	model_dirr   rt   ru   rv   ry   extra_session_options)r   pad_idF)add_special_tokensr   
max_lengthr#   )r   )	pad_tokenr   length)_load_onnx_modelr   r   r   rt   ru   rv   ry   r   r	   rX   r\   special_token_to_id
MASK_TOKENr   paddingr4   stringpunctuationencodeidsr3   
truncationenable_truncationenable_paddingrg   )r;   _symbolcurrent_max_lengths       rB   rb   zColbert.load_onnx_model   sr   oo--88LLnnnn"&"="= 	 	
 #14??"Ka~~)))!55dooF NN228< !,,
, NN!!&U!CGGJ,
 "^^66|D((4F4J(K..:Lq:P.Q++oo%%(( 	, 	

s   #1F0parallelc              +     K    | j                   d| j                  t        | j                        |||| j                  | j
                  | j                  | j                  | j                  | j                  d|E d{    y7 w)a  
        Encode a list of documents into list of embeddings.
        We use mean pooling with attention so that the model can handle variable-length inputs.

        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
        )rr   rs   rN   r^   r   ru   rv   rw   r|   rz   r   N )
_embed_documentsrr   rd   rs   ru   rv   rw   r   r   r   )r;   rN   r^   r   r(   s        rB   embedzColbert.embed   s}     , )4(( 
$..)!nn!33 $ 9 9"&"="=
 
 	
 	
s   A;B=B>Bc              +      K   t        |t              r|g}t        | d      r| j                  | j	                          |D ]/  }| j                  | j                  |gd      d      E d {    1 y 7 w)Nr   F)r'   )rc   rd   ra   r   rb   rC   
onnx_embed)r;   rP   r(   texts       rB   query_embedzColbert.query_embed  sv     eS!GEtW%);  "D55u5e 6    s   A)A5+A3,A5c                     t         S rW   )ColbertEmbeddingWorkerro   s    rB   _get_worker_classzColbert._get_worker_class!  s    %%rM   )T)i   TF)r)   N)   N)&__name__
__module____qualname__rG   rF   rg   r   r   boolr   r   r
   rC   dictrd   rL   listr   rU   rR   rQ   intrm   classmethodr   rq   r   AUTOr   r   r~   rb   r   r   r   r   r   __classcell__)r   s   @rB   r"   r"   (   s    J 9=5'5155HK5	*	54 AE
sJ/
9=
PS
	c:o	

$s) 
T 
C 
TXYaTb 
S T(^ 
T#Y 4>  "'Xc]"  	
    
@ (t,A'B ( ( !%"37$kk'+ $*.E#E# :E# t	E#
 L)D0E# VmE# I$E# E# :E# !4ZE# E#N
@ #	#
#&#
 #
 *	#

 #
 
*	#
J
x}!4 
 
Q[H\ 
 &$'::'F"G & &rM   r"   c                   $    e Zd ZdedededefdZy)r   rr   rs   r(   r)   c                      t        d||dd|S )Nr#   )rr   rs   rt   r   )r"   )r;   rr   rs   r(   s       rB   init_embeddingz%ColbertEmbeddingWorker.init_embedding'  s'     
!
 	
 	
rM   N)r   r   r   rd   r   r"   r   r   rM   rB   r   r   &  s$    
 
 
 
PW 
rM   r   )$r   typingr   r   r   r   numpyr5   
tokenizersr   r   #fastembed.common.preprocessor_utilsr	   fastembed.common.typesr
   r   fastembed.commonr   fastembed.common.onnx_modelr   fastembed.common.utilsr   r   :fastembed.late_interaction.late_interaction_embedding_baser   fastembed.text.onnx_text_modelr   r   "fastembed.common.model_descriptionr   r   r    r   __annotations__r"   r   r   rM   rB   <module>r      s     0 0  * > 5 ) 9 ? N Q &g78 5gFG'9 $45 ,{&.j0I {&|
0< 
rM   