Ë
    �…jH  ã                   óZ   — d dl mZ d dlmZ d dlmZ d dlmZ  ee«      Z	 G d„ de«      Z
y)é    )Ú	getLogger)ÚFusionAttentionClip)Ú
ModelProto)ÚBertOnnxModelc                   ó:   ‡ — e Zd Zddededefˆ fd„Zd„ Zd„ Zˆ xZS )ÚClipOnnxModelÚmodelÚ	num_headsÚhidden_sizec                 óv   •— t         ‰| �  |||¬«       t        | | j                  | j                  «      | _        y )N)r
   r   )ÚsuperÚ__init__r   r   r
   Úclip_attention_fusion)Úselfr	   r
   r   Ú	__class__s       €úi/root/aria/tools/markitdown-venv/lib/python3.12/site-packages/onnxruntime/transformers/onnx_model_clip.pyr   zClipOnnxModel.__init__   s5   ø€ Ü‰Ñ˜¨)ÀÐÔMÜ%8¸¸t×?OÑ?OÐQU×Q_ÑQ_Ó%`ˆÕ"ó    c                 óŽ   — i }g d¢}|D ]!  }| j                  |«      }t        |«      ||<   Œ# t        j                  d|› �«       |S )z8
        Returns node count of fused operators.
        )Ú	AttentionÚFastGeluÚGeluÚLayerNormalizationÚ	QuickGeluÚBiasGeluÚSkipLayerNormalizationzOptimized operators:)Úget_nodes_by_op_typeÚlenÚloggerÚinfo)r   Úop_countÚopsÚopÚnodess        r   Úget_fused_operator_statisticsz+ClipOnnxModel.get_fused_operator_statistics   sY   € ð ˆò
ˆð ò 	&ˆBØ×-Ñ-¨bÓ1ˆEÜ˜u›:ˆH�RŠLð	&ô 	�‰Ð*¨8¨*Ð5Ô6Øˆr   c                 ó8   — | j                   j                  «        y )N)r   Úapply)r   s    r   Úfuse_attentionzClipOnnxModel.fuse_attention)   s   € Ø×"Ñ"×(Ñ(Õ*r   )r   r   )	Ú__name__Ú
__module__Ú__qualname__r   Úintr   r$   r'   Ú__classcell__)r   s   @r   r   r      s+   ø„ ña˜jð a°Sð aÈ3õ aòö*+r   r   N)Úloggingr   Úfusion_attention_clipr   Úonnxr   Úonnx_model_bertr   r(   r   r   © r   r   ú<module>r2      s)   ðõ å 5Ý Ý )á	�8Ó	€ô+�Mõ +r   