Ë
    �…j=b  ã                   óz   — d dl mZ d dlZd dlmZ d dlmZ d dlm	Z	m
Z
mZmZ d dlmZ  ee«      Z G d„ de«      Zy)	é    )Ú	getLoggerN)ÚFusion)ÚFusionUtils)Ú	NodeProtoÚTensorProtoÚhelperÚnumpy_helper)Ú	OnnxModelc                   ó$  ‡ — e Zd ZdZdefˆ fd„Zd dedefd„Zdede	defd	„Z
d
ededefd„Zdededz  fd„Zdededz  fd„Zdedefd„Zdedeeef   de	fd„Zdededz  fd„Zdededz  fd„Zdededz  fd„Zdedededededefd„Zd„ Zˆ xZS )!ÚFusionMultiHeadAttentionMMDitzO
    Fuse MultiHeadAttention for Multimodal Diffusion Transformer (MMDiT).
    Úmodelc                 ó:   •— t         ‰| �  |ddg¬«       i | _        y )NÚMultiHeadAttentionÚSoftmax)Úfused_op_typeÚsearch_op_types)ÚsuperÚ__init__Úunsqueeze_update_map)Úselfr   Ú	__class__s     €új/root/aria/tools/markitdown-venv/lib/python3.12/site-packages/onnxruntime/transformers/fusion_mha_mmdit.pyr   z&FusionMultiHeadAttentionMMDit.__init__   s$   ø€ Ü‰Ñ˜Ð.BÐU^ÐT_ÐÔ`Ø$&ˆÕ!ó    Ú
start_nodeÚreturnc                 ó0  — | j                   j                  |g d¢|ddg|¬«      }|€y|d   }t        |j                  «      dk7  ry| j                   j	                  |j                  d   «      }|€yt        |j
                  «      dk7  ryt        |d   «      S )a�  
        Detect num_heads from Reshape & Transpose of q/k/v for both Stable Diffusion 3.x and Flux 1.x:

                MatMul    .. [-1] [24] ..
                 |        |  |  /   /
                Add     Concat(axis=0)
                  |      /
                  Reshape
                     |
                 Transpose(perm=0,1,3,2)
                     |
               (start_node)
        )Ú	TransposeÚReshapeÚConcatr   é   ©Úoutput_name_to_nodeéÿÿÿÿé   é   )r   Úmatch_parent_pathÚlenÚinputÚget_constant_valueÚshapeÚint)r   r   r"   Úinput_indexÚnodesÚconcat_shapeÚvalues          r   Úget_num_headsz+FusionMultiHeadAttentionMMDit.get_num_heads   s¡   € ð —
‘
×,Ñ,ØÒ:¸[È!ÈQÐ<OÐexð -ó 
ˆð ˆ=Øà˜R‘yˆÜˆ|×!Ñ!Ó" aÒ'Øà—
‘
×-Ñ-¨l×.@Ñ.@ÀÑ.CÓDˆØˆ=Øäˆu�{‰{Ó˜qÒ Øä�5˜‘8‹}Ðr   Útranspose_kÚconcat_before_transposec                 óî   — |r;| j                   j                  |ddgddg|¬«      }|r| j                  |d   |«      S y| j                   j                  |dgdg|¬«      }|r| j                  |d   |«      S y)aå  
                Detect num_heads from subgraph like the following (num_heads=24 in this example):
                               MatMu    .. [-1] [24] ..
                                 |       |  |  /   /
                                Add     Concat
                                  |      /
                                 Reshape
                                    |
                             Transpose(perm=0,2,1,3)
                                    |
                             SimplifiedLayerNormalization
                                    |
                            Transpose(perm=0,1,3,2)

                Another variant is to an extra Concat node to join two symmetrical subgraphs:

                           |              |
                          MatMul        MatMul   .. [-1] [24] ..
                           |              |       |  |  /   /
                          Add  Concat    Add      Concat
                            |  /          |      /
                          Reshape         Reshape
                            |              |
                         Transpose     Transpose(perm=0,2,1,3)
                            |              |
        SimplifiedLayerNormalization  SimplifiedLayerNormalization
                                |     /
                               Concat
                                 |
                            Transpose(perm=0,1,3,2)

                    Both patterns are used in stable diffusion 3.5 model.
        r   ÚSimplifiedLayerNormalizationr   r    r!   )r   r&   r0   )r   r1   r"   r2   r-   s        r   Úget_num_heads_from_kz2FusionMultiHeadAttentionMMDit.get_num_heads_from_k:   s¢   € ñD #Ø—J‘J×0Ñ0Ø˜hÐ(FÐGÈ!ÈQÈÐexð 1ó ˆEñ Ø×)Ñ)¨%°©(Ð4GÓHÐHð ð —J‘J×0Ñ0ØÐ<Ð=À¸sÐXkð 1ó ˆEñ Ø×)Ñ)¨%°©(Ð4GÓHÐHàr   Ú
input_nameÚoutput_namec                 óì  — d}| j                   j                  |«      }|€Tt        j                  t	        j
                  g d¢d¬«      |¬«      }| j                   j                  || j                  «       t        j                  d||g|g| j                   j                  d«      ¬«      }| j                  j                  |«       | j                  | j                  |j                  <   |j                  d   S )	a+  Add a Reshape node to convert 4D BxSxNxH to 3D BxSxD.

        Args:
            input_name (str): input name for the 4D tensor of shape BxSxNxH.
            output_name (str): output name for the 3D tensor of shape BxSxD, where D = N * H.

        Returns:
            str: the output name
        Úbsnh_to_bsd_reshape_dims)r   r   r#   Úint64)Údtype)Únamer   ©ÚinputsÚoutputsr<   r   )r   Úget_initializerr	   Ú
from_arrayÚnpÚarrayÚadd_initializerÚthis_graph_namer   Ú	make_nodeÚcreate_node_nameÚnodes_to_addÚappendÚnode_name_to_graph_namer<   Úoutput)r   r6   r7   Únew_dims_nameÚnew_dimsÚ	reshape_qs         r   Úreshape_to_3dz+FusionMultiHeadAttentionMMDit.reshape_to_3dk   sÎ   € ð 3ˆØ—:‘:×-Ñ-¨mÓ<ˆØÐÜ#×.Ñ.¬r¯x©xº
È'Ô/RÐYfÔgˆHØ�J‰J×&Ñ& x°×1EÑ1EÔFÜ×$Ñ$ØØ Ð.Ø �MØ—‘×,Ñ,¨YÓ7ô	
ˆ	ð 	×Ñ× Ñ  Ô+Ø7;×7KÑ7Kˆ×$Ñ$ Y§^¡^Ñ4Ø×Ñ Ñ"Ð"r   Úmul_qNc                 óF  — | j                   j                  |ddgddg«      }|€y|\  }}t        j                  |dg d¢«      sy|j                  d   |j                  d<   |j
                  d   }|dz   |j
                  d<   | j                  |j
                  d   |dz   «      S )	aÍ  
        MultiHeadAttenion requires query in BSD format. This function adjusts query from BNSH to BSD format.

        Before:
                               MatMul
                                 |
                               Add      Concat
                                 |      /
                                 Reshape
                                  |
                               Transpose(perm=0,2,1,3)
                                  |
                       SimplifiedLayerNorm
                                  |
                                 Mul

        After:
                               MatMul
                                 |
                                Add      Concat
                                 |      /
                                 Reshape
                                   |
                           SimplifiedLayerNorm
                                   |
                        Reshape (shape=[0, 0, -1])
        r4   r   r   NÚperm©r   r%   r    é   Ú_BSNHÚ_BSD)r   r&   r   Úcheck_node_attributer(   rK   rO   )r   rP   r"   ÚpathÚsln_aÚtranspose_aÚ
sln_outputs          r   Ú'adjust_query_from_bnsh_to_bsd_no_concatzEFusionMultiHeadAttentionMMDit.adjust_query_from_bnsh_to_bsd_no_concat…   s¬   € ð: �z‰z×+Ñ+ØØ+¨[Ð9Ø�ˆFó
ˆð
 ˆ<ØØ!Ñˆˆ{ä×/Ñ/°¸VÂ\ÔRØð %×*Ñ*¨1Ñ-ˆ�‰�A‰Ø—\‘\ !‘_ˆ
Ø$ wÑ.ˆ�‰�Q‰à×!Ñ! %§,¡,¨q¡/°:ÀÑ3FÓGÐGr   c                 ó|  — | j                   j                  |g d¢g d¢«      }|€y|\  }}}t        |j                  «      dk7  ry| j                   j                  |ddgddg«      }|€y|\  }}t	        j
                  |d	g d
¢«      syt	        j
                  |d	g d
¢«      syt	        j
                  |dd«      sy|j                  d   |j                  d<   |j                  d   |j                  d<   t        j                  d|j                  d   |j                  d   g|j                  d   dz   g| j                   j                  d«      d¬«      }	| j                  j                  |	«       | j                  | j                  |	j                  <   | j                  |	j                  d   |j                  d   dz   «      S )a•  
        MultiHeadAttenion requires query in BSD format. This function adjusts query from BNSH to BSD format.

            Before:
                      MatMul      MatMul
                        |            |
                        Add Concat  Add    Concat
                         |    /      |      /
                         Reshape     Reshape
                            |           |
        Transpose(perm=0,2,1,3)      Transpose(perm=0,2,1,3)
                            |           |
            SimplifiedLayerNorm  SimplifiedLayerNorm
                            |     /
                            Concat(axis=2)
                             |
                            Mul

            After:
                   MatMul        MatMul
                     |              |
                    Add Concat     Add     Concat
                     |    /         |     /
                     Reshape       Reshape
                        |            |
           SimplifiedLayerNorm  SimplifiedLayerNorm
                        |       /
                      Concat(axis=1)
                         |
                      Reshape (shape=[0, 0, -1])
        )r   r4   r   )r   r   r   Nr%   r4   r   r    r   rR   rS   Úaxisr   rU   ©r>   r?   r<   r^   rV   )r   r&   r'   r(   r   rW   r   rF   rK   rG   rH   rI   rE   rJ   r<   rO   )
r   rP   r"   rX   ÚconcatrY   rZ   Úsln_bÚtranspose_bÚnew_concat_nodes
             r   Úadjust_query_from_bnsh_to_bsdz;FusionMultiHeadAttentionMMDit.adjust_query_from_bnsh_to_bsdµ   sª  € ðB �z‰z×+Ñ+ØÚCÚó
ˆð
 ˆ<ØØ%)Ñ"ˆ��{äˆv�|‰|Ó Ò!Øà�z‰z×+Ñ+ØØ+¨[Ð9Ø�ˆFó
ˆð
 ˆ<ØØ!Ñˆˆ{ä×/Ñ/°¸VÂ\ÔRØä×/Ñ/°¸VÂ\ÔRØä×/Ñ/°¸ÀÔBØð %×*Ñ*¨1Ñ-ˆ�‰�A‰Ø$×*Ñ*¨1Ñ-ˆ�‰�A‰ä ×*Ñ*ØØ—L‘L ‘O U§\¡\°!¡_Ð5Ø—]‘] 1Ñ%¨Ñ/Ð0Ø—‘×,Ñ,¨XÓ6Øô
ˆð 	×Ñ× Ñ  Ô1Ø=A×=QÑ=Qˆ×$Ñ$ _×%9Ñ%9Ñ:à×!Ñ! /×"8Ñ"8¸Ñ";¸V¿]¹]È1Ñ=MÐPVÑ=VÓWÐWr   Ú	unsqueezec                 ón  — | j                   j                  |j                  «      }|�€Œt        |j                  «      dk(  rPt        j                  d|j                  |j                  d   dz   g| j                  j                  d«      dg¬«      }n¾d}| j                  j                  |«      €Ot        j                  |t        j                  dgdg¬«      }| j                  j                  || j                  «       t        j                  d|j                  d   |g|j                  d   dz   g| j                  j                  d«      ¬	«      }| j                   j#                  |«       | j                  | j$                  |j                  <   |j                  d   }|| j                   |j                  <   |S )
Nr    Ú	Unsqueezer   rU   r%   )r>   r?   r<   ÚaxesÚunsqueeze_axes_2)r<   Ú	data_typeÚdimsÚvalsr=   )r   Úgetr<   r'   r(   r   rF   rK   r   rG   r@   Úmake_tensorr   ÚINT64rD   rE   rH   rI   rJ   )r   re   Úupdated_unsqueeze_outputÚnew_nodeÚinitializer_nameri   s         r   Úupdate_unsqueeze_axes_1_to_2z:FusionMultiHeadAttentionMMDit.update_unsqueeze_axes_1_to_2  sƒ  € Ø#'×#<Ñ#<×#@Ñ#@ÀÇÁÓ#PÐ Ø#Ñ+Ü�9—?‘?Ó# qÒ(Ü!×+Ñ+ØØ$Ÿ?™?Ø&×-Ñ-¨aÑ0°7Ñ:Ð;ØŸ™×4Ñ4°[ÓAØ˜ô‘ð $6Ð Ø—:‘:×-Ñ-Ð.>Ó?ÐGÜ'-×'9Ñ'9Ø-Ü"-×"3Ñ"3Ø˜SØ˜Sô	(Ð$ð —J‘J×.Ñ.Ð/?À×AUÑAUÔVä!×+Ñ+ØØ%ŸO™O¨AÑ.Ð0@ÐAØ&×-Ñ-¨aÑ0°7Ñ:Ð;ØŸ™×4Ñ4°[ÓAô	�ð ×Ñ×$Ñ$ XÔ.Ø:>×:NÑ:NˆD×(Ñ(¨¯©Ñ7Ø'/§¡°qÑ'9Ð$Ø8PˆD×%Ñ% i§n¡nÑ5à'Ð'r   Úaddr"   c                 ól  — t        |j                  «      dk7  ry| j                  j                  |g d¢g d¢|«      }|€yt	        | j                  «      }|j                  |d   «      }|�|dgk7  ry|j                  |d   «      }|�|dgk7  ry| j                  j                  |g d¢g d¢|«      }|€y|j                  |d   «      }|�|dgk7  ry|j                  |d   «      }|�|dgk7  ry| j                  |d   «      |d   j                  d<   | j                  |d   «      |d   j                  d<   y)	a®  
        Update axes of Unsqueeze from [1] to [2] in the following pattern:
                  Unsqueeze        Unsqueeze
                  (axes=[0])       (axes=[0])
                     |              |
                  Unsqueeze        Unsqueeze
              ... (axes=[1])  ...  (axes=[1])
                |     /        |   /
                   Mul         Mul
                    |       /
                     Add
        Args:
            add (NodeProto): the Add node
            output_name_to_node (Dict[str, NodeProto]): mapping from output name to node

        Returns:
            bool: True if the pattern is matched and updated successfully, False otherwise.
        r%   F)ÚMulrg   rg   )r    r    r   r    r   )r   r    r   T)r'   r(   r   r&   r   Úget_squeeze_or_unsqueeze_axesrs   )r   rt   r"   Únodes_bÚfusion_utilsÚaxes_1Úaxes_0Únodes_as           r   Úupdate_unsqueeze_axesz3FusionMultiHeadAttentionMMDit.update_unsqueeze_axes(  sL  € ô& ˆs�y‰y‹>˜QÒØð —*‘*×.Ñ.¨sÒ4UÒW`ÐbuÓvˆØˆ?Øä" 4§:¡:Ó.ˆØ×;Ñ;¸GÀA¹JÓGˆØˆ>˜V¨ sš]Øà×;Ñ;¸GÀA¹JÓGˆØˆ>˜V¨ sš]Øð —*‘*×.Ñ.¨sÒ4UÒW`ÐbuÓvˆØˆ?Øà×;Ñ;¸GÀA¹JÓGˆØˆ>˜V¨ sš]Øà×;Ñ;¸GÀA¹JÓGˆØˆ>˜V¨ sš]Øà"×?Ñ?ÀÈÁ
ÓKˆ�‰
×Ñ˜ÑØ"×?Ñ?ÀÈÁ
ÓKˆ�‰
×Ñ˜ÑØr   c                 ó  — | j                   j                  |g d¢g d¢«      }|€y|\  }}}}}t        |j                  «      dk7  ry| j                   j                  |ddgddg«      }|€y|\  }	}
t	        j
                  |d	g d
¢«      syt	        j
                  |
d	g d
¢«      syt	        j
                  |dd«      sy| j                  ||«      sy|j                  d   |j                  d<   |
j                  d   |	j                  d<   t        j                  d|j                  d   |	j                  d   g|j                  d   dz   g| j                   j                  d«      d¬«      }| j                  j                  |«       | j                  | j                  |j                  <   | j                   j!                  |j                  d   |j                  d   «       | j#                  |j                  d   |j                  d   dz   «      S )a3  
        Adjust graph to change query format from BNSH to BSD for Flux model.
        Note that the graph pattern is complex, and we only do a shallow match here.

        Before:
                       |               |
        Transpose(perm=0,2,1,3)    Transpose(perm=0,2,1,3)
                        |              |
        SimplifiedLayerNorm  SimplifiedLayerNorm
                        |             /
                        Concat(axis=2)
                         |
                        Mul     Mul
                         |    /
                          Add
                           |
                          Mul

        After (Transpose nods are removed, and a Reshape is added):

                        |           |
            SimplifiedLayerNorm  SimplifiedLayerNorm
                        |         /
                    Concat(axis=1)
                        |
                        Mul    Mul
                         |    /
                          Add
                           |
                       Reshape (shape=[0, 0, -1])
        )ÚAddrv   r   r4   r   )r   r   r   r   r   Nr%   r4   r   r    r   rR   rS   r^   r   rU   r_   rV   )r   r&   r'   r(   r   rW   r}   r   rF   rK   rG   rH   rI   rE   rJ   r<   Úreplace_input_of_all_nodesrO   )r   rP   r"   rX   rt   Ú_mul_ar`   rY   rZ   ra   rb   rc   s               r   Ú"adjust_flux_query_from_bnsh_to_bsdz@FusionMultiHeadAttentionMMDit.adjust_flux_query_from_bnsh_to_bsd]  sè  € ðB �z‰z×+Ñ+ØÚQÚó
ˆð
 ˆ<ØØ26Ñ/ˆˆV�V˜U Käˆv�|‰|Ó Ò!Øà�z‰z×+Ñ+ØØ+¨[Ð9Ø�ˆFó
ˆð
 ˆ<ØØ!Ñˆˆ{ä×/Ñ/°¸VÂ\ÔRØä×/Ñ/°¸VÂ\ÔRØä×/Ñ/°¸ÀÔBØð ×)Ñ)¨#Ð/BÔCØð %×*Ñ*¨1Ñ-ˆ�‰�A‰Ø$×*Ñ*¨1Ñ-ˆ�‰�A‰ä ×*Ñ*ØØ—L‘L ‘O U§\¡\°!¡_Ð5Ø—]‘] 1Ñ%¨Ñ/Ð0Ø—‘×,Ñ,¨XÓ6Øô
ˆð 	×Ñ× Ñ  Ô1Ø=A×=QÑ=Qˆ×$Ñ$ _×%9Ñ%9Ñ:Ø�
‰
×-Ñ-¨f¯m©m¸AÑ.>À×@VÑ@VÐWXÑ@YÔZà×!Ñ! #§*¡*¨Q¡-°·±¸A±ÀÑ1GÓHÐHr   c                 ó†  — | j                   j                  |g d¢g d¢«      }|€y|\  }}}}t        j                  |dg d¢«      sy| j	                  ||«      sy|j
                  d   |j
                  d<   |j                  d   dz   |j                  d<   | j                  |j                  d   |j                  d   dz   «      S )	a0  
        Adjust graph to change query format from BNSH to BSD for Flux model.
        Note that the graph pattern is complex, and we only do a shallow match here.

        Before:
                      |
                    Transpose(perm=0,2,1,3)
                      |
                    SimplifiedLayerNorm
                      |
                     Mul     Mul
                       |   /
                       Add
                        |
                       Mul

        After (Transpose is removed, and a Reshape is added):

                        |
                      SimplifiedLayerNorm
                        |
                        Mul   Mul
                         |   /
                         Add
                          |
                       Reshape (shape=[0, 0, -1])
        )r   rv   r4   r   )r   r   r   r   NrR   rS   r   rU   rV   )r   r&   r   rW   r}   r(   rK   rO   )r   rP   r"   rX   rt   r�   rY   rZ   s           r   Ú)adjust_flux_single_query_from_bnsh_to_bsdzGFusionMultiHeadAttentionMMDit.adjust_flux_single_query_from_bnsh_to_bsd±  sÀ   € ð: �z‰z×+Ñ+ØÚGÚó
ˆð
 ˆ<ØØ*.Ñ'ˆˆV�U˜Kä×/Ñ/°¸VÂ\ÔRØð ×)Ñ)¨#Ð/BÔCØð %×*Ñ*¨1Ñ-ˆ�‰�A‰ØŸ
™
 1™¨Ñ/ˆ�
‰
�1‰à×!Ñ! #§*¡*¨Q¡-°·±¸A±ÀÑ1GÓHÐHr   Úqc           	      ó&  — t        j                  d|g|dz   g| j                  j                  dd¬«      g d¢¬«      }| j                  j                  |«       | j                  | j                  |j                  <   | j                  |dz   |dz   «      S )Nr   rU   ÚTranspose_BNSH_to_BSNH)Úname_prefixrS   )r<   rR   rV   )
r   rF   r   rG   rH   rI   rE   rJ   r<   rO   )r   r…   r"   Útranspose_qs       r   Útranspose_reshape_bnsh_to_bsdz;FusionMultiHeadAttentionMMDit.transpose_reshape_bnsh_to_bsdä  s‹   € Ü×&Ñ&ØØˆCØ�‰[ˆMØ—‘×,Ñ,¨[ÐF^Ð,Ó_Úô
ˆð 	×Ñ× Ñ  Ô-Ø9=×9MÑ9Mˆ×$Ñ$ [×%5Ñ%5Ñ6à×!Ñ! ! g¡+¨q°6©zÓ:Ð:r   ÚkÚvrK   Ú	num_headsc                 óö   — |dkD  sJ ‚|||g}|g}t        j                  d||| j                  j                  d«      ¬«      }d|_        |j
                  j                  t        j                  d|«      g«       |S )a~  
        Create a MultiHeadAttention node.

        Args:
            q (str): name of q
            k (str): name of k
            v (str): name of v
            output (str): output name of MHA
            num_heads (int): number of attention heads. If a model is pruned, it is the number of heads after pruning.

        Returns:
            NodeProto: the node created.
        r   r   r=   zcom.microsoftr�   )r   rF   r   rG   ÚdomainÚ	attributeÚextendÚmake_attribute)	r   r…   r‹   rŒ   rK   r�   Ú
mha_inputsÚmha_outputsÚmha_nodes	            r   Úcreate_multihead_attention_nodez=FusionMultiHeadAttentionMMDit.create_multihead_attention_nodeñ  s†   € ð, ˜1Š}Ðˆ}ð ˜˜A�Yˆ
ð �hˆä×#Ñ#Ø ØØØ—‘×,Ñ,Ð-AÓBô	
ˆð *ˆŒØ×Ñ×!Ñ!¤6×#8Ñ#8¸ÀiÓ#PÐ"QÔRð ˆr   c                 ó  — |j                   dk(  sJ ‚|}| j                  j                  |j                  d   «      ry | j                  j	                  |g d¢g d¢|«      }|€y |\  }}}t        j                  |dg d¢«      sy | j                  j                  |g d¢g d¢«      }	|	€y |	\  }
}}}}}}}|j                  d   }||j                  d   k7  ry | j                  j                  |
d	d
gddg«      }|€y |\  }}|j                  d   }t        j                  |dg d¢«      sy | j                  j                  |ddgddg«      }|€y |d   j                  d   |j                  d   k7  ry |j                  d   }| j                  j                  |dd|¬«      }|�x| j                  j                  |d
d|¬«      }|€y t        j                  |dg d¢«      sy | j                  j                  |d
d|¬«      }|€y t        j                  |dg d¢«      s=y | j                  j                  |d
d|¬«      }|€y t        j                  |dg d¢«      sy |r| j                  ||«      n| j                  ||d¬«      }|dk(  r| j                  |||d u«      }|dk  ry |�| j                  ||«      }n| j                  ||«      }|€:| j                  ||«      }|€&| j                  ||«      }|€| j!                  ||«      }| j#                  ||||j                  d   |¬«      }| j$                  j'                  |«       | j(                  | j*                  |j,                  <   | j.                  j1                  |||g«       d| _        y )Nr   r   )ÚMatMulr   r   )©r   r   r™   r™   rR   rS   )r˜   rv   ÚSqrtÚDivrš   ÚCastÚSliceÚShape)r   r   r    r   r    r   r   r   rv   r   r    )r   r    rT   r%   rš   r›   r   )r,   r"   )r,   )r…   r‹   rŒ   rK   r�   T)Úop_typer   Úfind_graph_outputrK   Úmatch_child_pathr   rW   r&   r(   Úmatch_parentr0   r5   rd   r\   r‚   r„   rŠ   r–   rH   rI   rE   rJ   r<   Únodes_to_remover‘   Úprune_graph)r   ÚnodeÚinput_name_to_nodesr"   Úsoftmaxr-   Ú
matmul_s_vÚtranspose_outÚreshape_outÚq_nodesÚ	matmul_qkrP   Úsqrt_q_2Údiv_qÚsqrt_qÚ_Úshape_qÚq_bnshÚk_nodesÚmul_kr1   r‹   Úk_scale_nodesrŒ   Úconcat_vÚtranspose_1Útranspose_2r�   Úqueryrq   s                                 r   Úfusez"FusionMultiHeadAttentionMMDit.fuse  sí  € Ø�|‰|˜yÒ(Ð(Ð(Øˆð �:‰:×'Ñ'¨¯©°qÑ(9Ô:Øà—
‘
×+Ñ+ØÒ7Ò9QÐSfó
ˆð ˆ=Øà16Ñ.ˆ
�M ;Ü×/Ñ/°¸vÂ|ÔTØà—*‘*×.Ñ.ØÚNÚ$ó
ˆð ˆ?ØàCJÑ@ˆ	�5˜( E¨6°1°a¸à—‘˜Q‘ˆØ�W—]‘] 1Ñ%Ò%Øà—*‘*×.Ñ.¨y¸5À+Ð:NÐQRÐTUÐPVÓWˆØˆ?Øà$Ñˆˆ{Ø×Ñ˜aÑ ˆÜ×/Ñ/°¸VÂ\ÔRØàŸ
™
×4Ñ4°U¸VÀU¸OÈaÐQRÈVÓTˆØÐ ØØ˜Ñ×!Ñ! !Ñ$¨¯©°qÑ(9Ò9Øà×Ñ˜QÑˆð —:‘:×*Ñ*¨:°xÈQÐdwÐ*ÓxˆØÐð Ÿ*™*×1Ñ1Ø˜+°1ÐJ]ð 2ó ˆKð Ð"ØÜ×3Ñ3°KÀÊÔVØàŸ*™*×1Ñ1Ø˜+°1ÐJ]ð 2ó ˆKð Ð"ØÜ×3Ñ3°KÀÊÔVØð Ÿ*™*×1Ñ1Ø˜K°QÐL_ð 2ó ˆKð Ð"ØÜ×3Ñ3°KÀÊÔVØñ
 ð ×Ñ˜xÐ)<Ô=à×#Ñ# JÐ0CÐQRÐ#ÓSð 	ð ˜Š>à×1Ñ1°+Ð?RÐT\ÐdhÐThÓiˆIØ˜AŠ~Øð ÐØ×6Ñ6°uÐ>QÓR‰Eà×@Ñ@ÀÐH[Ó\ˆEàˆ=Ø×;Ñ;¸EÐCVÓWˆEØˆ}Ø×FÑFÀuÐNaÓb�Ø�=ð !×>Ñ>¸vÐGZÓ[�Eà×7Ñ7ØØØØ×%Ñ% aÑ(Øð 8ó 
ˆð 	×Ñ× Ñ  Ô*Ø6:×6JÑ6Jˆ×$Ñ$ X§]¡]Ñ3à×Ñ×#Ñ# Z°ÀÐ$LÔMð  ˆÕr   )r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r
   r   r   r+   r0   Úboolr5   ÚstrrO   r\   rd   rs   Údictr}   r‚   r„   rŠ   r–   rº   Ú__classcell__)r   s   @r   r   r      sv  ø„ ñð'˜iõ 'ñ¨	ð ÐZ]ó ðB/°	ð /Ðimð /Ðruó /ðb#¨ð #¸#ð #À#ó #ð4.H¸Yð .HÐ`cÐfjÑ`jó .Hð`MX°9ð MXÐVYÐ\`ÑV`ó MXð^"(°ið "(ÀCó "(ðH3¨ð 3ÈÈcÐS\ÈnÑI]ð 3Ðbfó 3ðjRI¸	ð RIÐ[^ÐaeÑ[eó RIðh1I¸yð 1IÐbeÐhlÑbló 1Iðf;¨sð ;ÈCÐRVÉJó ;ð)àð)ð ð)ð ð	)ð
 ð)ð ð)ð 
ó)öV r   r   )Úloggingr   ÚnumpyrB   Úfusion_baser   ry   r   Úonnxr   r   r   r	   Ú
onnx_modelr
   r»   Úloggerr   © r   r   ú<module>rÊ      s4   ðõ
 ã Ý Ý $ß =Ó =Ý  á	�8Ó	€ôK
  Fõ K
 r   