
    ^Nj-                         d dl Z d dlmZ d dlmZ d dlmZmZmZm	Z	m
Z
mZ d dlZd dlZd dlmZ d dlmZ d dlmZmZmZ d dlmZ  ed	      Ze G d
 d             Z G d dee         Z G d deee         Zy)    N)	dataclass)Path)AnyGenericIterableSequenceTypeTypeVar)NDArray)	Tokenizer)OnnxProvider
NumpyArrayDevice)WorkerTc                       e Zd ZU eed<   dZeej                     dz  ed<   dZ	eej                     dz  ed<   dZ
eeef   dz  ed<   y)OnnxOutputContextmodel_outputNattention_mask	input_idsmetadata)__name__
__module____qualname__r   __annotations__r   r   npint64r   r   dictstrr        l/Users/ahmed/devFolder/Ultron/claude-voice/.venv/lib/python3.12/site-packages/fastembed/common/onnx_model.pyr   r      sO    /3NGBHH%,3*.Iwrxx 4'.&*Hd38nt#*r!   r   c                   p   e Zd ZdZeded   fd       Zdedede	e
   fdZdd	Zd
eeef   dedeeef   fdZdej$                  ddfdedededz  dee   dz  deez  dedz  deeef   dz  ddfdZedeeef   deeef   fd       Zedej6                  deeef   ddfd       ZddZdededefdZy)	OnnxModel)enable_cpu_mem_arenareturnEmbeddingWorker[T]c                     t        d      N%Subclasses must implement this methodNotImplementedError)clss    r"   _get_worker_classzOnnxModel._get_worker_class   s    !"IJJr!   outputkwargsc                     t        d      )ah  Post-process the ONNX model output to convert it into a usable format.

        Args:
            output (OnnxOutputContext): The raw output from the ONNX model.
            **kwargs: Additional keyword arguments that may be needed by specific implementations.

        Returns:
            Iterable[T]: Post-processed output as an iterable of type T.
        r*   r+   )selfr/   r0   s      r"   _post_process_onnx_outputz#OnnxModel._post_process_onnx_output"   s     ""IJJr!   Nc                      d | _         d | _        y N)model	tokenizerr2   s    r"   __init__zOnnxModel.__init__.   s    26
+/r!   
onnx_inputc                     |S )z,
        Preprocess the onnx input.
        r    )r2   r:   r0   s      r"   _preprocess_onnx_inputz OnnxModel._preprocess_onnx_input2   s
     r!   	model_dir
model_filethreads	providerscuda	device_idextra_session_optionsc                 P   ||z  }t        j                         }	d|	v }
|du xs |t        j                  k(  }|r%|#t	        j
                  d| d| dt        d       |t        |      }n(|s|t        j                  k(  r|
r|dg}ndd|ifg}nd	g}g }|D ]?  }t        |t              r|n|d
   }|j                  |       ||	vs0t        d| d|	        t        j                         }t         j                  j                  |_        |||_        ||_        || j'                  ||       t        j(                  t        |      ||      | _        d|v rL| j*                  J | j*                  j-                         }d|vrt	        j
                  d| dt.               y y y )NCUDAExecutionProviderTz@`cuda` and `providers` are mutually exclusive parameters, cuda: z, providers: zY. If you'd like to use providers, cuda should be one of [False, Device.CPU, Device.AUTO].   )category
stacklevelrB   CPUExecutionProviderr   z	Provider z( is not available. Available providers: )r@   sess_optionsz@Attempt to set CUDAExecutionProvider failed. Current providers: z.If you are using CUDA 12.x, install onnxruntime-gpu via `pip install onnxruntime-gpu --extra-index-url https://aiinfra.pkgs.visualstudio.com/PublicPackages/_packaging/onnxruntime-cuda-12/pypi/simple/`)ortget_available_providersr   CUDAwarningswarnUserWarninglistAUTO
isinstancer   append
ValueErrorSessionOptionsGraphOptimizationLevelORT_ENABLE_ALLgraph_optimization_levelintra_op_num_threadsinter_op_num_threadsadd_extra_session_optionsInferenceSessionr6   get_providersRuntimeWarning)r2   r=   r>   r?   r@   rA   rB   rC   
model_pathavailable_providerscuda_availableexplicit_cudaonnx_providersrequested_provider_namesproviderprovider_namesocurrent_providerss                     r"   _load_onnx_modelzOnnxModel._load_onnx_model:   s    +
!99;04GG;(;Y2MMmI; 745 %  !)_Ntv{{2~ "9!:#:[)<T"U!V45N.0 &H(28S(AHxPQ{M$++M:$77 .VWjVkl  ' !&)&@&@&O&O#&-B#&-B# ,**2/DE))
O~B

 #&>>::))) $

 8 8 :&.??VWhVi jg g #	 @ ?r!   model_kwargsc                 t    |j                         D ci c]  \  }}|| j                  v s|| c}}S c c}}w )zA convenience method to select the exposed session options in models

        Args:
            model_kwargs (dict[str, Any]): The model kwargs.

        Returns:
            dict[str, Any]: a dict with filtered exposed session options.
        )itemsEXPOSED_SESSION_OPTIONS)r-   rk   kvs       r"   _select_exposed_session_optionsz)OnnxModel._select_exposed_session_options   s<     ".!3!3!5Z!5Ac>Y>Y9Y1!5ZZZs   44session_optionsextra_optionsc                 z    |D ]'  }|| j                   v rJ | d| j                    d        d|v r|d   |_        yy)aC  Add extra session options to the existing options object in-place

        Args:
            session_options (ort.SessionOptions): The existing session options object.
            extra_options (dict[str, Any]): The extra session options available in cls.EXPOSED_SESSION_OPTIONS.

        Returns:
            None
        z- is unknown or not exposed (exposed options: )r%   N)rn   r%   )r-   rr   rs   options       r"   r\   z#OnnxModel.add_extra_session_options   sa     $F#555fFsGbGbFccdef5 $ "]23@AW3XO0 3r!   c                     t        d      r)   r+   r8   s    r"   load_onnx_modelzOnnxModel.load_onnx_model       !"IJJr!   argsc                     t        d      r)   r+   )r2   rz   r0   s      r"   
onnx_embedzOnnxModel.onnx_embed   ry   r!   )r&   N)r   r   r   rn   classmethodr	   r.   r   r   r   r   r3   r9   r   r   r   r<   r   rR   r   intr   r   boolrj   rq   rK   rV   r\   rx   r|   r    r!   r"   r$   r$      s   7K$';"< K K
K0A 
KS 
KU]^_U` 
K0sJ/;>	c:o	 48$kk $7;CC C t	C
 L)D0C VmC :C  $CH~4C 
CJ 	[4S> 	[dSVX[S[n 	[ 	[ Y!00YAEc3hY	Y Y&KK Ks K7H Kr!   r$   c            	           e Zd Zdedededee   fdZdededefdZe	dedededdfd       Z
d	eeeef      deeeef      fd
Zy)EmbeddingWorker
model_name	cache_dirr0   r&   c                     t               r5   r+   r2   r   r   r0   s       r"   init_embeddingzEmbeddingWorker.init_embedding   s     "##r!   c                 6     | j                   ||fi || _        y r5   )r   r6   r   s       r"   r9   zEmbeddingWorker.__init__   s     )T((YI&I
r!   r'   c                      | d||d|S )N)r   r   r    r    )r-   r   r   r0   s       r"   startzEmbeddingWorker.start   s    HjIHHHr!   rm   c                     t        d      r)   r+   )r2   rm   s     r"   processzEmbeddingWorker.process   ry   r!   N)r   r   r   r   r   r$   r   r   r9   r}   r   r   tupler~   r   r    r!   r"   r   r      s    $$ $ 	$
 
1$JJ J 	J Is Is Ic IFZ I IKXeCHo6 K8E#s(O;T Kr!   r   )rN   dataclassesr   pathlibr   typingr   r   r   r   r	   r
   numpyr   onnxruntimerK   numpy.typingr   
tokenizersr   fastembed.common.typesr   r   r   fastembed.parallel_processorr   r   r   r$   r   r    r!   r"   <module>r      sv     !  B B       C C / CL + + +HK
 HKVKfgaj Kr!   