§
    kŠtj=b  ã                   óˆ   — 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dS )	é    )Ú	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 | _        d S )NÚMultiHeadAttentionÚSoftmax)Úfused_op_typeÚsearch_op_types)ÚsuperÚ__init__Úunsqueeze_update_map)Úselfr   Ú	__class__s     €úg/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/onnxruntime/transformers/fusion_mha_mmdit.pyr   z&FusionMultiHeadAttentionMMDit.__init__   s5   ø€ Ý‰Œ×Ò˜Ð.BÐU^ÐT_ÐÑ`Ô`Ð`Ø$&ˆÔ!Ð!Ð!ó    r   Ú
start_nodeÚreturnc                 óD  — | j                              |g d¢|ddg|¬¦  «        }|€dS |d         }t          |j        ¦  «        dk    rdS | j                              |j        d         ¦  «        }|€dS t          |j        ¦  «        dk    rdS 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_nodeNéÿÿÿÿé   é   )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ð -ñ 
ô 
ˆð ˆ=Ø�1à˜R”yˆÝˆ|Ô!Ñ"Ô" aÒ'Ð'Ø�1à”
×-Ò-¨lÔ.@ÀÔ.CÑDÔDˆØˆ=Ø�1åˆuŒ{ÑÔ˜qÒ Ð Ø�1å�5˜”8‰}Œ}Ðr   Útranspose_kÚconcat_before_transposec                 ó  — |rA| j                              |ddgddg|¬¦  «        }|r|                      |d         |¦  «        S n>| j                              |dgdg|¬¦  «        }|r|                      |d         |¦  «        S dS )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 #ð 	IØ”J×0Ò0Ø˜hÐ(FÐGÈ!ÈQÈÐexð 1ñ ô ˆEð ð IØ×)Ò)¨%°¬(Ð4GÑHÔHÐHðIð ”J×0Ò0ØÐ<Ð=À¸sÐXkð 1ñ ô ˆEð ð IØ×)Ò)¨%°¬(Ð4GÑHÔHÐHàˆqr   Ú
input_nameÚoutput_namec                 ó²  — d}| j                              |¦  «        }|€Lt          j        t	          j        g d¢d¬¦  «        |¬¦  «        }| j                              || j        ¦  «         t          j	        d||g|g| j          
                    d¦  «        ¬¦  «        }| 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_dimsN)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¸
¸
¸
È'Ð/RÑ/RÔ/RÐYfÐgÑgÔgˆHØŒJ×&Ò& x°Ô1EÑFÔFÐFÝÔ$ØØ Ð.Ø �MØ”×,Ò,¨YÑ7Ô7ð	
ñ 
ô 
ˆ	ð 	Ô× Ò  Ñ+Ô+Ð+Ø7;Ô7KˆÔ$ Y¤^Ñ4ØÔ Ô"Ð"r   Úmul_qNc                 ó.  — | j                              |ddgddg¦  «        }|€dS |\  }}t          j        |dg d¢¦  «        sdS |j        d         |j        d<   |j        d         }|dz   |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
                                 |
                               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×+Ò+ØØ+¨[Ð9Ø�ˆFñ
ô 
ˆð
 ˆ<Ø�4Ø!Ñˆˆ{åÔ/°¸VÀ\À\À\ÑRÔRð 	Ø�4ð %Ô*¨1Ô-ˆŒ�A‰Ø”\ !”_ˆ
Ø$ wÑ.ˆŒ�Q‰à×!Ò! %¤,¨q¤/°:ÀÑ3FÑGÔGÐGr   c                 ó2  — | j                              |g d¢g d¢¦  «        }|€dS |\  }}}t          |j        ¦  «        dk    rdS | j                              |ddgddg¦  «        }|€dS |\  }}t	          j        |d	g d
¢¦  «        sdS t	          j        |d	g d
¢¦  «        sdS t	          j        |dd¦  «        sdS |j        d         |j        d<   |j        d         |j        d<   t          j        d|j        d         |j        d         g|j        d         dz   g| j          	                    d¦  «        d¬¦  «        }	| 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×+Ò+ØØCÐCÐCØˆIˆIñ
ô 
ˆð
 ˆ<Ø�4Ø%)Ñ"ˆ��{åˆvŒ|ÑÔ Ò!Ð!Ø�4àŒz×+Ò+ØØ+¨[Ð9Ø�ˆFñ
ô 
ˆð
 ˆ<Ø�4Ø!Ñˆˆ{åÔ/°¸VÀ\À\À\ÑRÔRð 	Ø�4åÔ/°¸VÀ\À\À\ÑRÔRð 	Ø�4åÔ/°¸ÀÑBÔBð 	Ø�4ð %Ô*¨1Ô-ˆŒ�A‰Ø$Ô*¨1Ô-ˆŒ�A‰å Ô*ØØ”L ”O U¤\°!¤_Ð5Ø”] 1Ô%¨Ñ/Ð0Ø”×,Ò,¨XÑ6Ô6Øð
ñ 
ô 
ˆð 	Ô× Ò  Ñ1Ô1Ð1Ø=AÔ=QˆÔ$ _Ô%9Ñ:à×!Ò! /Ô"8¸Ô";¸V¼]È1Ô=MÐPVÑ=VÑWÔWÐWr   Ú	unsqueezec                 óô  — | j                              |j        ¦  «        }|�€Ut          |j        ¦  «        dk    rGt          j        d|j        |j        d         dz   g| j         	                    d¦  «        dg¬¦  «        }n¬d}| j         
                    |¦  «        €Dt          j        |t          j        dgdg¬¦  «        }| j                             || j        ¦  «         t          j        d|j        d         |g|j        d         dz   g| j         	                    d¦  «        ¬	¦  «        }| 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Ô#PÐ Ø#Ñ+Ý�9”?Ñ#Ô# qÒ(Ð(Ý!Ô+ØØ$œ?Ø&Ô-¨aÔ0°7Ñ:Ð;Øœ×4Ò4°[ÑAÔAØ˜ðñ ô ��ð $6Ð Ø”:×-Ò-Ð.>Ñ?Ô?ÐGÝ'-Ô'9Ø-Ý"-Ô"3Ø˜SØ˜Sð	(ñ (ô (Ð$ð ”J×.Ò.Ð/?ÀÔAUÑVÔVÐVå!Ô+ØØ%œO¨AÔ.Ð0@ÐAØ&Ô-¨aÔ0°7Ñ:Ð;Øœ×4Ò4°[ÑAÔAð	ñ ô �ð Ô×$Ò$ XÑ.Ô.Ð.Ø:>Ô:NˆDÔ(¨¬Ñ7Ø'/¤°qÔ'9Ð$Ø8PˆDÔ% i¤nÑ5à'Ð'r   Úaddr"   c                 óÊ  — t          |j        ¦  «        dk    rdS | j                             |g d¢g d¢|¦  «        }|€dS t	          | j        ¦  «        }|                     |d         ¦  «        }|�|dgk    rdS |                     |d         ¦  «        }|�|dgk    rdS | j                             |g d¢g d¢|¦  «        }|€dS |                     |d         ¦  «        }|�|dgk    rdS |                     |d         ¦  «        }|�|dgk    rdS |                      |d         ¦  «        |d         j        d<   |                      |d         ¦  «        |d         j        d<   d	S )
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   Nr    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(  sŽ  € õ& ˆsŒy‰>Œ>˜QÒÐØ�5ð ”*×.Ò.¨sÐ4UÐ4UÐ4UÐW`ÐW`ÐW`ÐbuÑvÔvˆØˆ?Ø�5å" 4¤:Ñ.Ô.ˆØ×;Ò;¸GÀA¼JÑGÔGˆØˆ>˜V¨ sš]˜]Ø�5à×;Ò;¸GÀA¼JÑGÔGˆØˆ>˜V¨ sš]˜]Ø�5ð ”*×.Ò.¨sÐ4UÐ4UÐ4UÐW`ÐW`ÐW`ÐbuÑvÔvˆØˆ?Ø�5à×;Ò;¸GÀA¼JÑGÔGˆØˆ>˜V¨ sš]˜]Ø�5à×;Ò;¸GÀA¼JÑGÔGˆØˆ>˜V¨ sš]˜]Ø�5à"×?Ò?ÀÈÄ
ÑKÔKˆ�Œ
Ô˜ÑØ"×?Ò?ÀÈÄ
ÑKÔKˆ�Œ
Ô˜ÑØˆtr   c                 óÈ  — | j                              |g d¢g d¢¦  «        }|€dS |\  }}}}}t          |j        ¦  «        dk    rdS | j                              |ddgddg¦  «        }|€dS |\  }	}
t	          j        |d	g d
¢¦  «        sdS t	          j        |
d	g d
¢¦  «        sdS t	          j        |dd¦  «        sdS |                      ||¦  «        sdS |j        d         |j        d<   |
j        d         |	j        d<   t          j        d|j	        d         |	j	        d         g|j	        d         dz   g| j          
                    d¦  «        d¬¦  «        }| j                             |¦  «         | j        | j        |j        <   | j                              |j	        d         |j	        d         ¦  «         |                      |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×+Ò+ØØQÐQÐQØˆOˆOñ
ô 
ˆð
 ˆ<Ø�4Ø26Ñ/ˆˆV�V˜U KåˆvŒ|ÑÔ Ò!Ð!Ø�4àŒz×+Ò+ØØ+¨[Ð9Ø�ˆFñ
ô 
ˆð
 ˆ<Ø�4Ø!Ñˆˆ{åÔ/°¸VÀ\À\À\ÑRÔRð 	Ø�4åÔ/°¸VÀ\À\À\ÑRÔRð 	Ø�4åÔ/°¸ÀÑBÔBð 	Ø�4ð ×)Ò)¨#Ð/BÑCÔCð 	Ø�4ð %Ô*¨1Ô-ˆŒ�A‰Ø$Ô*¨1Ô-ˆŒ�A‰å Ô*ØØ”L ”O U¤\°!¤_Ð5Ø”] 1Ô%¨Ñ/Ð0Ø”×,Ò,¨XÑ6Ô6Øð
ñ 
ô 
ˆð 	Ô× Ò  Ñ1Ô1Ð1Ø=AÔ=QˆÔ$ _Ô%9Ñ:ØŒ
×-Ò-¨f¬m¸AÔ.>ÀÔ@VÐWXÔ@YÑZÔZÐZà×!Ò! #¤*¨Q¤-°´¸A´ÀÑ1GÑHÔHÐHr   c                 ót  — | j                              |g d¢g d¢¦  «        }|€dS |\  }}}}t          j        |dg d¢¦  «        sdS |                      ||¦  «        sdS |j        d         |j        d<   |j        d         dz   |j        d<   |                      |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×+Ò+ØØGÐGÐGØˆLˆLñ
ô 
ˆð
 ˆ<Ø�4Ø*.Ñ'ˆˆV�U˜KåÔ/°¸VÀ\À\À\ÑRÔRð 	Ø�4ð ×)Ò)¨#Ð/BÑCÔCð 	Ø�4ð %Ô*¨1Ô-ˆŒ�A‰Øœ
 1œ¨Ñ/ˆŒ
�1‰à×!Ò! #¤*¨Q¤-°´¸A´ÀÑ1GÑHÔHÐHr   Úqc           	      ó
  — t          j        d|g|dz   g| j                             dd¬¦  «        g d¢¬¦  «        }| 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ˆÔ$ [Ô%5Ñ6à×!Ò! ! g¡+¨q°6©zÑ:Ô:Ð:r   ÚkÚvrK   Ú	num_headsc                 óð   — |dk    sJ ‚|||g}|g}t          j        d||| j                             d¦  «        ¬¦  «        }d|_        |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ÔBð	
ñ 
ô 
ˆð *ˆŒØÔ×!Ò!¥6Ô#8¸ÀiÑ#PÔ#PÐ"QÑRÔRÐRð ˆr   c                 ó  — |j         dk    sJ ‚|}| j                             |j        d         ¦  «        rd S | j                             |g d¢g d¢|¦  «        }|€d S |\  }}}t          j        |dg d¢¦  «        sd S | j                             |g d¢g d¢¦  «        }	|	€d S |	\  }
}}}}}}}|j        d         }||j        d         k    rd S | j                             |
d	d
gddg¦  «        }|€d S |\  }}|j        d         }t          j        |dg d¢¦  «        sd S | j                             |ddgddg¦  «        }|€d S |d         j        d         |j        d         k    rd S |j        d         }| j         	                    |dd|¬¦  «        }|�y| j         	                    |d
d|¬¦  «        }|€d S t          j        |dg d¢¦  «        sd S | j         	                    |d
d|¬¦  «        }|€d S t          j        |dg d¢¦  «        sd S n<| j         	                    |d
d|¬¦  «        }|€d S t          j        |dg d¢¦  «        sd S |r|  
                    ||¦  «        n|  
                    ||d¬¦  «        }|dk    r!|                      |||d u¦  «        }|dk    rd S |�|                      ||¦  «        }n|                      ||¦  «        }|€F|                      ||¦  «        }|€.|                      ||¦  «        }|€|                      ||¦  «        }|                      ||||j        d         |¬¦  «        }| j                             |¦  «         | j        | j        |j        <   | j                             |||g¦  «         d| _        d S )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Ñ:Ô:ð 	ØˆFà”
×+Ò+ØÐ7Ð7Ð7Ð9QÐ9QÐ9QÐSfñ
ô 
ˆð ˆ=ØˆFà16Ñ.ˆ
�M ;ÝÔ/°¸vÀ|À|À|ÑTÔTð 	ØˆFà”*×.Ò.ØØNÐNÐNØ$Ð$Ð$ñ
ô 
ˆð ˆ?ØˆFàCJÑ@ˆ	�5˜( E¨6°1°a¸à”˜Q”ˆØ�W”] 1Ô%Ò%Ð%ØˆFà”*×.Ò.¨y¸5À+Ð:NÐQRÐTUÐPVÑWÔWˆØˆ?ØˆFà$Ñˆˆ{ØÔ˜aÔ ˆÝÔ/°¸VÀ\À\À\ÑRÔRð 	ØˆFàœ
×4Ò4°U¸VÀU¸OÈaÐQRÈVÑTÔTˆØÐ ØˆFØ˜ÔÔ! !Ô$¨¬°qÔ(9Ò9Ð9ØˆFàÔ˜QÔˆð ”:×*Ò*¨:°xÈQÐdwÐ*ÑxÔxˆØÐð œ*×1Ò1Ø˜+°1ÐJ]ð 2ñ ô ˆKð Ð"Ø�ÝÔ3°KÀÈÈÈÑVÔVð Ø�àœ*×1Ò1Ø˜+°1ÐJ]ð 2ñ ô ˆKð Ð"Ø�ÝÔ3°KÀÈÈÈÑVÔVð Ø�ðð
 œ*×1Ò1Ø˜K°QÐL_ð 2ñ ô ˆKð Ð"Ø�ÝÔ3°KÀÈÈÈÑVÔVð Ø�ð
 ðTˆD×Ò˜xÐ)<Ñ=Ô=Ð=à×#Ò# JÐ0CÐQRÐ#ÑSÔSð 	ð ˜Š>ˆ>à×1Ò1°+Ð?RÐT\ÐdhÐThÑiÔiˆIØ˜AŠ~ˆ~Ø�ð ÐØ×6Ò6°uÐ>QÑRÔRˆEˆEà×@Ò@ÀÐH[Ñ\Ô\ˆEàˆ=Ø×;Ò;¸EÐCVÑWÔWˆEØˆ}Ø×FÒFÀuÐNaÑbÔb�Ø�=ð !×>Ò>¸vÐGZÑ[Ô[�Eà×7Ò7ØØØØÔ% aÔ(Øð 8ñ 
ô 
ˆð 	Ô× Ò  Ñ*Ô*Ð*Ø6:Ô6JˆÔ$ X¤]Ñ3àÔ×#Ò# Z°ÀÐ$LÑMÔMÐ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      s|  ø€ € € € € ðð ð'˜ið 'ð 'ð 'ð 'ð 'ð 'ðð ¨	ð ÐZ]ð ð ð ð ðB/°	ð /Ðimð /Ðruð /ð /ð /ð /ðb#¨ð #¸#ð #À#ð #ð #ð #ð #ð4.H¸Yð .HÐ`cÐfjÑ`jð .Hð .Hð .Hð .Hð`MX°9ð MXÐVYÐ\`ÑV`ð MXð MXð MXð MXð^"(°ið "(ÀCð "(ð "(ð "(ð "(ðH3¨ð 3ÈÈcÐS\ÈnÔI]ð 3Ðbfð 3ð 3ð 3ð 3ðjRI¸	ð RIÐ[^ÐaeÑ[eð RIð RIð RIð RIðh1I¸yð 1IÐbeÐhlÑblð 1Ið 1Ið 1Ið 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Ê      sÌ   ðð
 Ð Ð Ð Ð Ð à Ð Ð Ð Ø Ð Ð Ð Ð Ð Ø $Ð $Ð $Ð $Ð $Ð $Ø =Ð =Ð =Ð =Ð =Ð =Ð =Ð =Ð =Ð =Ð =Ð =Ø  Ð  Ð  Ð  Ð  Ð  à	ˆ�8Ñ	Ô	€ðK
 ð K
 ð K
 ð K
 ð K
  Fñ K
 ô K
 ð K
 ð K
 ð K
 r   