Ë
    �…j  ã                   óŽ   — d dl Z d dlmZ d dlmZ d dlmZ d dlmZ d dl	m
Z
 d dlmZ  e j                  e«      Z G d„ d	e«      Zy)
é    N)ÚFusionLayerNormalization)ÚFusionMultiHeadAttentionMMDit)ÚFusionOptions)Úis_installed)Ú
ModelProto)ÚBertOnnxModelc                   ór   ‡ — e Zd Zddededefˆ fd„Zd„ Zd„ Zd„ Zdd	e	dz  d
e
fd„Zdd	e	dz  fd„Zd„ Zˆ xZS )ÚMmditOnnxModelÚmodelÚ	num_headsÚhidden_sizec                 ó\   •— |dk(  r|dk(  s|dkD  r||z  dk(  sJ ‚t         ‰| �  |||¬«       y)ak  Initialize Multimodal Diffusion Transformer (MMDiT) ONNX Model.

        Args:
            model (ModelProto): the ONNX model
            num_heads (int, optional): number of attention heads. Defaults to 0 (detect the parameter automatically).
            hidden_size (int, optional): hidden dimension. Defaults to 0 (detect the parameter automatically).
        r   )r   r   N)ÚsuperÚ__init__)Úselfr   r   r   Ú	__class__s       €új/root/aria/tools/markitdown-venv/lib/python3.12/site-packages/onnxruntime/transformers/onnx_model_mmdit.pyr   zMmditOnnxModel.__init__   sA   ø€ ð ˜Q’ ;°!Ò#3¸ÀQºÈ;ÐYbÑKbÐfgÒKgÐhÐhÜ‰Ñ˜¨)ÀÐÕMó    c                 óD   — | j                  «        | j                  «        y ©N)Úprune_graphÚremove_unused_constant)r   s    r   ÚpostprocesszMmditOnnxModel.postprocess   s   € Ø×ÑÔØ×#Ñ#Õ%r   c                 óp   — d}t         j                  d«       t        | | d¬«      }|j                  «        y )NTzwThe optimized model requires LayerNormalization with broadcast support. Please use onnxruntime-gpu>=1.21 for inference.)Úcheck_constant_and_dimensionÚforce)ÚloggerÚwarningr   Úapply)r   Úlayernorm_support_broadcastÚfusions      r   Úfuse_layer_normzMmditOnnxModel.fuse_layer_norm"   s<   € Ø&*Ð#Ü�‰ð>ô	
ô *ØÐ3NÐ/NÐVZô
ˆð 	�‰�r   c                 ó:   — t        | «      }|j                  «        y r   )r   r   )r   r!   s     r   Úfuse_multi_head_attentionz(MmditOnnxModel.fuse_multi_head_attention-   s   € Ü.¨tÓ4ˆØ�‰�r   NÚoptionsÚadd_dynamic_axesc                 ó   — |rJ ‚t        d«      rLdd l}ddlm}  |«       5  d}|j                  t	        |«      dd¬«      }| j                  ||«       d d d «       y t        j                  d«       | j                  |d «       y # 1 sw Y   y xY w)NÚtqdmr   )Úlogging_redirect_tqdmé   r!   )ÚinitialÚdescz<tqdm is not installed. Run optimization without progress bar)r   r(   Útqdm.contrib.loggingr)   ÚrangeÚ	_optimizer   Úinfo)r   r%   r&   r(   r)   ÚstepsÚprogress_bars          r   ÚoptimizezMmditOnnxModel.optimize1   s�   € Ù#Ð#Ð#ä˜ÔÛÝBá&Ó(ñ 6Ø�Ø#Ÿy™y¬¨u«¸qÀx˜yÓP�Ø—‘˜w¨Ô5÷6ð 6ô
 �K‰KÐVÔWØ�N‰N˜7 DÕ)÷6ð 6ús   ¡2BÂBc                 ór  — |�|j                   s| j                  «        | j                  j                  «        |r|j	                  d«       |�|j
                  r | j                  «        | j                  «        |r|j	                  d«       |�|j                  r| j                  «        |r|j	                  d«       |�|j                  r| j                  «        |r|j	                  d«       | j                  «        |r|j	                  d«       t        j                  d| j                  «       › �«       y )Né   zopset version: )Úenable_shape_inferenceÚdisable_shape_inferenceÚutilsÚremove_useless_cast_nodesÚupdateÚenable_layer_normr"   Úfuse_simplified_layer_normÚenable_geluÚ	fuse_geluÚenable_attentionr$   r   r   r0   Úget_opset_version)r   r%   r2   s      r   r/   zMmditOnnxModel._optimize@   s  € ØÐ¨×)GÒ)GØ×(Ñ(Ô*ð 	�
‰
×,Ñ,Ô.ÙØ×Ñ Ô"àˆO × 9Ò 9Ø× Ñ Ô"Ø×+Ñ+Ô-ÙØ×Ñ Ô"àˆO × 3Ò 3Ø�N‰NÔÙØ×Ñ Ô"àˆO × 8Ò 8Ø×*Ñ*Ô,ÙØ×Ñ Ô"à×ÑÔÙØ×Ñ Ô"ä�‰�o d×&<Ñ&<Ó&>Ð%?Ð@ÕAr   c                 óŽ   — i }g d¢}|D ]!  }| j                  |«      }t        |«      ||<   Œ# t        j                  d|› �«       |S )z8
        Returns node count of fused operators.
        )ÚFastGeluÚMultiHeadAttentionÚLayerNormalizationÚSimplifiedLayerNormalizationzOptimized operators:)Úget_nodes_by_op_typeÚlenr   r0   )r   Úop_countÚopsÚopÚnodess        r   Úget_fused_operator_statisticsz,MmditOnnxModel.get_fused_operator_statistics_   sY   € ð ˆò
ˆð ò 	&ˆBØ×-Ñ-¨bÓ1ˆEÜ˜u›:ˆH�RŠLð	&ô 	�‰Ð*¨8¨*Ð5Ô6Øˆr   )r   r   )NF)NN)Ú__name__Ú
__module__Ú__qualname__r   Úintr   r   r"   r$   r   Úboolr3   r/   rL   Ú__classcell__)r   s   @r   r
   r
      s`   ø„ ñ	N˜jð 	N°Sð 	NÈ3õ 	Nò&ò	òñ* °Ñ 4ð *Ètó *ñB °Ñ!5ó Bö>r   r
   )ÚloggingÚfusion_layernormr   Úfusion_mha_mmditr   Úfusion_optionsr   Úimport_utilsr   Úonnxr   Úonnx_model_bertr   Ú	getLoggerrM   r   r
   © r   r   ú<module>r\      s<   ðó å 5Ý :Ý (Ý %Ý Ý )à	ˆ×	Ñ	˜8Ó	$€ô^�]õ ^r   