Ë
    �…jî  ã                   ó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)ÚFusion)ÚNumpyHelper)Ú	OnnxModelc                   ó@   ‡ — e Zd Zdefˆ fd„Zˆ fd„Zd„ Zd„ Zd„ Zˆ xZ	S )ÚFusionConstantFoldÚmodelc                 ó8   •— t         ‰| �  |ddg«       d| _        y )NÚ Ú	Transposer   )ÚsuperÚ__init__Úcount)Úselfr	   Ú	__class__s     €ún/root/aria/tools/markitdown-venv/lib/python3.12/site-packages/onnxruntime/transformers/fusion_constant_fold.pyr   zFusionConstantFold.__init__   s   ø€ Ü‰Ñ˜  [ MÔ2Øˆ�
ó    c                 ó†   •— t         ‰| �  «        | j                  dkD  r#t        j	                  d| j                  › �«       y y )Nr   zConstant Folded: )r   Úapplyr   ÚloggerÚinfo)r   r   s    €r   r   zFusionConstantFold.apply   s4   ø€ Ü‰‰ŒØ�:‰:˜Š>Ü�K‰KÐ+¨D¯J©J¨<Ð8Õ9ð r   c                 óP   — | j                  |||«       | j                  |||«       y)zX
        Apply multiple fusions on Transpose nodes that can be constant folded.
        N)Úfuse_1Úfuse_2)r   ÚnodeÚinput_name_to_nodesÚoutput_name_to_nodes       r   ÚfusezFusionConstantFold.fuse   s(   € ð 	�‰�DÐ-Ð/BÔCØ�‰�DÐ-Ð/BÕCr   c                 óÐ  — t        |j                  «      dk7  st        |j                  «      dk7  rt        j	                  d«       y| j
                  j                  |j                  d   «      }|€t        j	                  d«       yd}||j                  d      D ]-  }|j                  dk(  rt        |j                  «      dk(  rŒ+d} n |rt        j	                  d	«       y||j                  d      D ]%  }|j                  d
k(  rŒ|j                  dk(  rŒ#d} n |rt        j	                  d«       yt        j                  |«      }t        |j                  «      dk7  rt        j	                  d«       y|j                  }|j                  }	| j                  |«       | j                  ||	|j                  d   |j                  d   g|j                  ¬«       ||j                  d      D ]Æ  }t!        t        |j                  «      «      D ]£  }
|j                  |
   |j                  d   k(  sŒ#|j                  d   |j                  |
<   |j                  d
k(  sŒO|
dk(  s|
dk(  sŒZ|
dk(  rdnd}t#        |j$                  «      D ])  \  }}|j                  |k(  sŒd|j$                  |   _        Œ+ Œ¥ ŒÈ | j(                  j+                  |«       | xj,                  dz  c_        y)z¶
        Constant fold any initializer data representing a MatMul's
        weights that are stored in a Transpose op

        Ex: Transpose --> Gemm or Transpose --> MatMul
        é   ú:fuse_constant_fold: node has more than one input or outputNr   z8fuse_constant_fold: failed to identify initializer inputFr   TzAfuse_constant_fold: other non-Transpose nodes use the initializerÚGemmÚMatMulzOfuse_constant_fold: other non-Gemm and non-MatMul nodes use the transposed dataé   z7fuse_constant_fold: shape of initializer data is not 2D)ÚnameÚ	data_typeÚdimsÚvalsÚtransAÚtransB)ÚlenÚinputÚoutputr   Údebugr	   Úget_initializerÚop_typer   Úto_arrayÚshaper%   r&   Úremove_initializerÚadd_initializerÚTÚrangeÚ	enumerateÚ	attributeÚiÚnodes_to_removeÚappendr   )r   r   r   r   ÚprotoÚskipÚ
child_nodeÚweightr%   Údtyper9   ÚkeyÚjÚattr_keys                 r   r   zFusionConstantFold.fuse_1    sž  € ô ˆt�z‰z‹?˜aÒ¤3 t§{¡{Ó#3°qÒ#8Ü�L‰LÐUÔVØð —
‘
×*Ñ*¨4¯:©:°a©=Ó9ˆØˆ=Ü�L‰LÐSÔTØð ˆØ-¨d¯j©j¸©mÑ<ò 	ˆJØ×&Ñ&¨+Ò5¼#¸d¿j¹j»/ÈQÓ:NØ�Ùð	ñ Ü�L‰LÐ\Ô]Øð .¨d¯k©k¸!©nÑ=ò 	ˆJØ×&Ñ&¨&Ó0°J×4FÑ4FÈ(Ó4RØ�Ùð	ñ Ü�L‰LÐjÔkØô ×%Ñ% eÓ,ˆÜˆv�|‰|Ó Ò!Ü�L‰LÐRÔSØð �z‰zˆØ—‘ˆØ×Ñ Ô&Ø×ÑØØØ—,‘,˜q‘/ 6§<¡<°¡?Ð3Ø—‘ð	 	ô 	
ð .¨d¯k©k¸!©nÑ=ò 
	>ˆJÜœ3˜z×/Ñ/Ó0Ó1ò 	>�Ø×#Ñ# AÑ&¨$¯+©+°a©.Ó8Ø*.¯*©*°Q©-�J×$Ñ$ QÑ'à!×)Ñ)¨VÓ3¸¸aºÀ1ÈÃ6à*+¨qª&™h°h˜Ü+4°Z×5IÑ5IÓ+Jò >™K˜A˜xØ'Ÿ}™}°Ó3Ø<= 
× 4Ñ 4°QÑ 7Õ 9ñ>ñ	>ð
	>ð 	×Ñ×#Ñ# DÔ)Ø�
Š
�a‰Ž
r   c                 ó„  — t        |j                  «      dk7  st        |j                  «      dk7  rt        j	                  d«       y| j
                  j                  |dd«      }|€t        j	                  d«       yt        |j                  «      dk7  st        |j                  «      dk7  rt        j	                  d«       y|j                  d   j                  }|j                  d   j                  }||k7  rt        j	                  d«       y|j                  d   }||j                  d      }|D ]A  }	t        |	j                  «      D ]'  \  }
}||j                  d   k(  sŒ||	j                  |
<   Œ) ŒC | j                  j                  |«       | j                  j                  |«       | xj                  dz  c_        y)	zÎ
        Constant fold any Transpose --> Transpose ops since the root input
        is the final result

        Ex: root_input --> Transpose --> Transpose --> next_node to root_input --> next_node
        r    r!   Nr   r   z<fuse_constant_fold: failed to identify parent Transpose nodezAfuse_constant_fold: parent node has more than one input or outputz@fuse_constant_fold: Transpose node permutations aren't identical)r+   r,   r-   r   r.   r	   Úmatch_parentr8   Úintsr7   r:   r;   r   )r   r   r   r   Úparent_nodeÚ	node_permÚparent_node_permÚ
root_inputÚoutput_nodesÚoutput_noder9   Úinput_s               r   r   zFusionConstantFold.fuse_2h   s…  € ô ˆt�z‰z‹?˜aÒ¤3 t§{¡{Ó#3°qÒ#8Ü�L‰LÐUÔVØð —j‘j×-Ñ-¨d°KÀÓCˆØÐÜ�L‰LÐWÔXØÜˆ{× Ñ Ó! QÒ&¬#¨k×.@Ñ.@Ó*AÀQÒ*FÜ�L‰LÐ\Ô]Øà—N‘N 1Ñ%×*Ñ*ˆ	Ø&×0Ñ0°Ñ3×8Ñ8ÐàÐ(Ò(Ü�L‰LÐ[Ô\Øð !×&Ñ& qÑ)ˆ
Ø*¨4¯;©;°q©>Ñ:ˆØ'ò 	6ˆKÜ& {×'8Ñ'8Ó9ò 6‘	��6Ø˜TŸ[™[¨™^Ó+Ø+5�K×%Ñ% aÒ(ñ6ð	6ð 	×Ñ×#Ñ# DÔ)Ø×Ñ×#Ñ# KÔ0Ø�
Š
�a‰Ž
r   )
Ú__name__Ú
__module__Ú__qualname__r   r   r   r   r   r   Ú__classcell__)r   s   @r   r   r      s&   ø„ ð˜iõ ô:ò
DòFöP(r   r   N)Úloggingr   Úfusion_baser   Úfusion_utilsr   Ú
onnx_modelr   rN   r   r   © r   r   ú<module>rW      s+   ðõ å Ý $Ý  á	�8Ó	€ôA˜õ Ar   