Ë
    ÝÍ:jµ  ã                   ó„   — d dl m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e«      Z G d„ de«      Z G d„ d	e«      Zy
)é    )Ú	getLogger)ÚFusion)ÚFusionUtils)Ú	NodeProtoÚTensorProtoÚhelper)Ú	OnnxModelc                   óR   ‡ — e Zd Zdefˆ fd„Zdedeeee   f   deeef   fd„Z	ˆ xZ
S )ÚFusionTransposeÚmodelc                 ó(   •— t         ‰| �  |dd«       y )NÚ	Transpose©ÚsuperÚ__init__©Úselfr   Ú	__class__s     €ú~/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/onnxruntime/transformers/fusion_transpose.pyr   zFusionTranspose.__init__   s   ø€ Ü‰Ñ˜ ¨[Õ9ó    Útranspose_nodeÚinput_name_to_nodesÚoutput_name_to_nodec                 óÚ  — |}|j                   d   |vry||j                   d      }|j                  dk7  rd}nS|}| j                  j                  ||«      }|rt	        |«      dkD  ry|j                   d   |vry||j                   d      }|j                  dk7  ryt        j                  |d«      }t        |t        «      sJ ‚t        j                  |d«      }	t        |	t        «      sJ ‚t	        |	«      t	        |«      k(  sJ ‚g }
t        |«      D ]  \  }}|
j                  |	|   «       Œ |€>t        j                  | j                  |||«      rY| j                  j                  |«       n=t        j                  | j                  |||«      r| j                  j                  |«       |j                  d«       |j                  j!                  t#        j$                  d|
«      g«       y)a  
        Note that onnxruntime will do comprehensive transpose optimization after loading model.
        The purpose of this fusion is to make graph clean before running onnxruntime.

        Case 1:
              (input)-->Transpose(perm=a)-->Transpose(perm=b)-->
        After:
              (input)-->Transpose(perm=a)-->  (this path can be removed if the output is not used anymore)
                |
                +----->Transpose(perm=a*b)-->

        Case 2 (Cast has only one child):
              (input)-->Transpose(perm=a)--> Cast -->Transpose(perm=b)-->
        After:
              (input)-->Transpose(perm=a)-->  (this path can be removed if the output is not used anymore)
                |
                +----->Cast --> Transpose(perm=a*b)-->
        r   NÚCasté   r   ÚpermÚ	attribute)ÚinputÚop_typer   Úget_childrenÚlenr	   Úget_node_attributeÚ
isinstanceÚlistÚ	enumerateÚappendr   Úskip_parentÚnodes_to_removeÚ
ClearFieldr   Úextendr   Úmake_attribute)r   r   r   r   Útranspose_bÚtranspose_aÚ	cast_nodeÚcast_childrenÚpermutationÚparent_permutationÚoutput_permutationÚ_jÚindexs                r   ÚfusezFusionTranspose.fuse   s×  € ð0 %ˆØ×Ñ˜QÑÐ':Ñ:Øà)¨+×*;Ñ*;¸AÑ*>Ñ?ˆØ×Ñ &Ò(Ø‰Ià#ˆIà ŸJ™J×3Ñ3°IÐ?RÓSˆMÙ¤ ]Ó!3°aÒ!7Øà�‰˜qÑ!Ð)<Ñ<Øà-¨i¯o©o¸aÑ.@ÑAˆKà×Ñ +Ò-Øä×2Ñ2°;ÀÓGˆÜ˜+¤tÔ,Ð,Ð,ä&×9Ñ9¸+ÀvÓNÐÜÐ,¬dÔ3Ð3Ð3äÐ%Ó&¬#¨kÓ*:Ò:Ð:Ð:àÐÜ" ;Ó/ò 	A‰IˆB�Ø×%Ñ%Ð&8¸Ñ&?Õ@ð	Að ÐÜ×&Ñ& t§z¡z°;ÀÐM`ÔaØ×$Ñ$×+Ñ+¨KÕ8ä×&Ñ& t§z¡z°9¸kÐK^Ô_Ø×$Ñ$×+Ñ+¨KÔ8Ø×Ñ˜{Ô+Ø×Ñ×$Ñ$¤f×&;Ñ&;¸FÐDVÓ&WÐ%XÕYr   )Ú__name__Ú
__module__Ú__qualname__r	   r   r   ÚdictÚstrr%   r6   Ú__classcell__©r   s   @r   r   r      sO   ø„ ð:˜iõ :ðAZà!ðAZð " # t¨I¡Ð"6Ñ7ðAZð " # y .Ñ1÷	AZr   r   c                   ój   ‡ — e Zd Zdefˆ fd„Zd
dedee   fd„Zde	de
eee	   f   de
ee	f   fd	„Zˆ xZS )ÚFusionInsertTransposer   c                 ó(   •— t         ‰| �  |dd«       y )NÚ Ú	GroupNormr   r   s     €r   r   zFusionInsertTranspose.__init__Y   s   ø€ Ü‰Ñ˜  KÕ0r   Ú
input_namer   c                 óì   — | j                   j                  d«      }|€|dz   dz   |z   }t        j                  d|g|g|¬«      }|j                  j                  t        j                  d|«      g«       |S )z&Append a Transpose node after an inputr   Ú_outú-)ÚinputsÚoutputsÚnamer   )r   Úcreate_node_namer   Ú	make_noder   r+   r,   )r   rC   r   Úoutput_nameÚ	node_namer   s         r   Úcreate_transpose_nodez+FusionInsertTranspose.create_transpose_node\   sw   € à—J‘J×/Ñ/°Ó<ˆ	ØÐØ# fÑ,¨sÑ2°ZÑ?ˆKÜ×)Ñ)¨+¸z¸lÐU`ÐTaÐhqÔrˆØ× Ñ ×'Ñ'¬×)>Ñ)>¸vÀtÓ)LÐ(MÔNØÐr   Úgroup_norm_noder   r   c                 ó”  — | j                   j                  |g d¢g d¢«      }|€y|\  }}}}}	| j                   j                  |j                  d   «      ryt	        j
                  |d«      }
t        |
t        «      sJ ‚|
g d¢k7  ryt        |j                  «      dk(  ræ| j                   j                  |j                  d   «      d	k(  r»t        |j                  «      dk(  r£| j                   j                  |j                  d   «      dk(  rxt        | j                   j                  |	|«      «      dk(  rPt        | j                   j                  ||«      «      dk(  r(t        | j                   j                  ||«      «      dk(  syd
}| j                   j                  |«      €&| j                  |t        j                  dgdgd¬«       d}| j                   j                  |«      €&| j                  |t        j                  dgdgd¬«       d|j                  d<   d
|j                  d<   | j                   j!                  d«      dz   }| j                   j#                  |j                  d   |«       | j%                  |j                  d   g d¢|«      }| j                   j'                  || j(                  «       | j+                  d«       y)a  
        This optimization will insert an Transpose, and onnxruntime transpose optimizer will remove it together with
        another Transpose so that we can get effect of reducing one Transpose after onnxruntime optimization.
        Before:
            --> Gemm --> Unsqueeze(axes=[2]) --> Unsqueeze(axes=[3]) --> Add --> Transpose([0,2,3,1]) --> GroupNorm
        After:
            --> Gemm --> Unsqueeze(axes=[1]) --> Unsqueeze(axes=[2]) -->Transpose([0,3,1,2]) --> Add --> Transpose([0,2,3,1]) --> GroupNorm
        )r   ÚAddÚ	UnsqueezerR   ÚGemm)r   r   Nr   r   Nr   r   )r   é   é   r   rT   r   rU   Úort_const_unsqueeze_axes_1F)rI   Ú	data_typeÚdimsÚvalsÚrawÚort_const_unsqueeze_axes_2r   Ú_NCHW)r   rU   r   rT   zInsert Transpose)r   Úmatch_parent_pathÚfind_graph_outputÚoutputr	   r#   r$   r%   r"   r   Úget_constant_valuer!   Úget_initializerÚadd_initializerr   ÚINT64rJ   Úreplace_input_of_all_nodesrN   Úadd_nodeÚthis_graph_nameÚincrease_counter)r   rO   r   r   Ú	gemm_pathÚ	transposeÚaddÚunsqueeze_3Úunsqueeze_2Úgemmr1   Úaxes_1Úaxes_2Útranspose_output_nameÚnew_transposes                  r   r6   zFusionInsertTranspose.fusee   sŒ  € ð —J‘J×0Ñ0ØÒSÒUgó
ˆ	ð ÐØØ9BÑ6ˆ	�3˜ [°$Ø�:‰:×'Ñ'¨×(:Ñ(:¸1Ñ(=Ô>Øä×2Ñ2°9¸fÓEˆÜ˜+¤tÔ,Ð,Ð,Øš,Ò&Øô �×!Ñ!Ó" aÒ'Ø—
‘
×-Ñ-¨k×.?Ñ.?ÀÑ.BÓCÀqÒHÜ�K×%Ñ%Ó&¨!Ò+Ø—
‘
×-Ñ-¨k×.?Ñ.?ÀÑ.BÓCÀqÒHÜ�D—J‘J×+Ñ+¨DÐ2EÓFÓGÈ1ÒLÜ�D—J‘J×+Ñ+¨KÐ9LÓMÓNÐRSÒSÜ�D—J‘J×+Ñ+¨KÐ9LÓMÓNÐRSÒSàð .ˆØ�:‰:×%Ñ% fÓ-Ð5Ø× Ñ ØÜ%×+Ñ+Ø�SØ�SØð !ô ð .ˆØ�:‰:×%Ñ% fÓ-Ð5Ø× Ñ ØÜ%×+Ñ+Ø�SØ�SØð !ô ð  <ˆ×Ñ˜!ÑØ;ˆ×Ñ˜!ÑØ $§
¡
× ;Ñ ;¸KÓ HÈ7Ñ RÐØ�
‰
×-Ñ-¨k×.@Ñ.@ÀÑ.CÐEZÔ[Ø×2Ñ2°;×3EÑ3EÀaÑ3HÊ,ÐXmÓnˆØ�
‰
×Ñ˜M¨4×+?Ñ+?Ô@Ø×ÑÐ0Õ1r   )N)r7   r8   r9   r	   r   r;   r%   ÚintrN   r   r:   r6   r<   r=   s   @r   r?   r?   X   sb   ø„ ð1˜iõ 1ñ°ð ¸4À¹9ó ðB2à"ðB2ð " # t¨I¡Ð"6Ñ7ðB2ð " # y .Ñ1÷	B2r   r?   N)Úloggingr   Úfusion_baser   Úfusion_utilsr   Úonnxr   r   r   Ú
onnx_modelr	   r7   Úloggerr   r?   © r   r   ú<module>rz      sB   ðõ å Ý $ß /Ñ /Ý  á	�8Ó	€ôEZ�fô EZôPO2˜Fõ O2r   