§
    kŠtjî  ã                   óh   — 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
dS )é    )Ú	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                 ó^   •— t          ¦   «                              |ddg¦  «         d| _        d S )NÚ Ú	Transposer   )ÚsuperÚ__init__Úcount)Úselfr	   Ú	__class__s     €úk/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/onnxruntime/transformers/fusion_constant_fold.pyr   zFusionConstantFold.__init__   s,   ø€ Ý‰Œ×Ò˜  [ MÑ2Ô2Ð2ØˆŒ
ˆ
ˆ
ó    c                 ó¦   •— t          ¦   «                              ¦   «          | j        dk    r$t                               d| j        › �¦  «         d S d S )Nr   zConstant Folded: )r   Úapplyr   ÚloggerÚinfo)r   r   s    €r   r   zFusionConstantFold.apply   sI   ø€ Ý‰Œ�Š‰ŒˆØŒ:˜Š>ˆ>Ý�KŠKÐ8¨D¬JÐ8Ð8Ñ9Ô9Ð9Ð9Ð9ð ˆ>r   c                 ób   — |                       |||¦  «         |                      |||¦  «         dS )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ÔCÐCØ�Š�DÐ-Ð/BÑCÔCÐCÐCÐCr   c                 ó–  — t          |j        ¦  «        dk    st          |j        ¦  «        dk    rt                               d¦  «         dS | j                             |j        d         ¦  «        }|€t                               d¦  «         dS d}||j        d                  D ])}|j        dk    rt          |j        ¦  «        dk    sd} nŒ*|rt                               d	¦  «         dS ||j        d                  D ]}|j        d
k    s|j        dk    sd} nŒ|rt                               d¦  «         dS t          j	        |¦  «        }t          |j
        ¦  «        dk    rt                               d¦  «         dS |j        }|j        }	|                      |¦  «         |                      ||	|j
        d         |j
        d         g|j        ¬¦  «         ||j        d                  D ]¯}t!          t          |j        ¦  «        ¦  «        D ]‹}
|j        |
         |j        d         k    rm|j        d         |j        |
<   |j        d
k    rM|
dk    s|
dk    rA|
dk    rdnd}t#          |j        ¦  «        D ]"\  }}|j        |k    rd|j        |         _        Œ#ŒŒŒ°| j                             |¦  «         | xj        dz  c_        dS )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‰?Œ?˜aÒÐ¥3 t¤{Ñ#3Ô#3°qÒ#8Ð#8Ý�LŠLÐUÑVÔVÐVØˆFð ”
×*Ò*¨4¬:°a¬=Ñ9Ô9ˆØˆ=Ý�LŠLÐSÑTÔTÐTØˆFð ˆØ-¨d¬j¸¬mÔ<ð 	ð 	ˆJØÔ&¨+Ò5Ð5½#¸d¼j¹/¼/ÈQÒ:NÐ:NØ�Ø�ð ;Oð ð 	Ý�LŠLÐ\Ñ]Ô]Ð]ØˆFð .¨d¬k¸!¬nÔ=ð 	ð 	ˆJØÔ&¨&Ò0Ð0°JÔ4FÈ(Ò4RÐ4RØ�Ø�øØð 	Ý�LŠLÐjÑkÔkÐkØˆFõ Ô% eÑ,Ô,ˆÝˆvŒ|ÑÔ Ò!Ð!Ý�LŠLÐRÑSÔSÐSØˆFð ŒzˆØ”ˆØ×Ò Ñ&Ô&Ð&Ø×ÒØØØ”,˜q”/ 6¤<°¤?Ð3Ø”ð	 	ñ 	
ô 	
ð 	
ð .¨d¬k¸!¬nÔ=ð 
	>ð 
	>ˆJÝ�3˜zÔ/Ñ0Ô0Ñ1Ô1ð 	>ð 	>�ØÔ# AÔ&¨$¬+°a¬.Ò8Ð8Ø*.¬*°Q¬-�JÔ$ QÑ'à!Ô)¨VÒ3Ð3¸¸aº¸À1ÈÂ6À6à*+¨qª&¨&˜h˜h°h˜Ý+4°ZÔ5IÑ+JÔ+Jð >ð >™K˜A˜xØ'œ}°Ò3Ð3Ø<= 
Ô 4°QÔ 7Ô 9øøð	>ð 	Ô×#Ò# DÑ)Ô)Ð)Øˆ
Œ
�a‰ˆ
Œ
ˆ
ˆ
r   c                 ór  — t          |j        ¦  «        dk    st          |j        ¦  «        dk    rt                               d¦  «         dS | j                             |dd¦  «        }|€t                               d¦  «         dS t          |j        ¦  «        dk    st          |j        ¦  «        dk    rt                               d¦  «         dS |j        d         j        }|j        d         j        }||k    rt                               d¦  «         dS |j        d         }||j        d                  }|D ]7}	t          |	j        ¦  «        D ] \  }
}||j        d         k    r
||	j        |
<   Œ!Œ8| j
                             |¦  «         | j
                             |¦  «         | xj        dz  c_        dS )	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‰?Œ?˜aÒÐ¥3 t¤{Ñ#3Ô#3°qÒ#8Ð#8Ý�LŠLÐUÑVÔVÐVØˆFð ”j×-Ò-¨d°KÀÑCÔCˆØÐÝ�LŠLÐWÑXÔXÐXØˆFÝˆ{Ô Ñ!Ô! QÒ&Ð&­#¨kÔ.@Ñ*AÔ*AÀQÒ*FÐ*FÝ�LŠLÐ\Ñ]Ô]Ð]ØˆFà”N 1Ô%Ô*ˆ	Ø&Ô0°Ô3Ô8ÐàÐ(Ò(Ð(Ý�LŠLÐ[Ñ\Ô\Ð\ØˆFð !Ô& qÔ)ˆ
Ø*¨4¬;°q¬>Ô:ˆØ'ð 	6ð 	6ˆKÝ& {Ô'8Ñ9Ô9ð 6ð 6‘	��6Ø˜Tœ[¨œ^Ò+Ð+Ø+5�KÔ% aÑ(øð6ð
 	Ô×#Ò# DÑ)Ô)Ð)ØÔ×#Ò# KÑ0Ô0Ð0Øˆ
Œ
�a‰ˆ
Œ
ˆ
ˆ
r   )
Ú__name__Ú
__module__Ú__qualname__r   r   r   r   r   r   Ú__classcell__)r   s   @r   r   r      s‘   ø€ € € € € ð˜ið ð ð ð ð ð ð:ð :ð :ð :ð :ð
Dð Dð DðFð Fð FðP(ð (ð (ð (ð (ð (ð (r   r   N)Úloggingr   Úfusion_baser   Úfusion_utilsr   Ú
onnx_modelr   rN   r   r   © r   r   ú<module>rW      sœ   ðð Ð Ð Ð Ð Ð à Ð Ð Ð Ð Ð Ø $Ð $Ð $Ð $Ð $Ð $Ø  Ð  Ð  Ð  Ð  Ð  à	ˆ�8Ñ	Ô	€ðAð Að Að Að A˜ñ Aô Að Að Að Ar   