§
    kŠtjŒ  ã                   ó˜   — 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d
S )é    )Ú	getLogger)ÚFusion)ÚFusionUtils)Ú	NodeProtoÚTensorProtoÚhelper)Ú	OnnxModelc                   ó(  ‡ — e Zd ZdZd"dedefˆ fd„Zdeddeeef         z  fd	„Z	d
ede
eee         f         dedefd„Zd„ Zd„ Zd„ Zd„ Zd„ Zdedeedez  f         fd„Z	 	 	 d#ded
edededdez  dedz  fd„Zd„ Zd„ Z	 d$d„Zd„ Zd „ Zd!„ Zˆ xZS )%ÚFusionEmbedLayerNoMaskzŒ
    Fuse embedding layer into one node (EmbedLayerNormalization).
    It supports the following model types: BERT, DistilBert, ALBert.
    úno maskÚmodelÚdescriptionc                 ó´   •— t          ¦   «                              |dddg|¦  «         t          |¦  «        | _        d | _        d| _        d | _        d | _        d S )NÚEmbedLayerNormalizationÚLayerNormalizationÚSkipLayerNormalizationF)ÚsuperÚ__init__r   ÚutilsÚshape_inferÚshape_infer_doneÚ	attentionÚ
embed_node)Úselfr   r   Ú	__class__s      €úh/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/onnxruntime/transformers/fusion_embedlayer.pyr   zFusionEmbedLayerNoMask.__init__   sf   ø€ Ý‰Œ×ÒØØ%Ø!Ð#;Ð<Øñ		
ô 	
ð 	
õ ! Ñ'Ô'ˆŒ
ØˆÔØ %ˆÔð ˆŒØˆŒˆˆó    ÚaddÚreturnNc                 óª   — | j                              |dgdg¦  «        }|€d S | j                              |dgdg¦  «        }|€d S |d         |d         fS )NÚGatherr   é   )r   Úmatch_parent_path)r   r   Úgather_0_pathÚgather_1_paths       r   Úmatch_two_gatherz'FusionEmbedLayerNoMask.match_two_gather%   sg   € Øœ
×4Ò4°S¸8¸*ÀqÀcÑJÔJˆØÐ Ø�4àœ
×4Ò4°S¸8¸*ÀqÀcÑJÔJˆØÐ Ø�4à˜QÔ ¨qÔ!1Ð1Ð1r   Ú	layernormÚinput_name_to_nodesÚis_distil_bertc                 ó$  — | j                              |d|d¬¦  «        | _        | j        �dS |j        d         |vrdS ||j        d                  }t	          d„ |D ¦   «         ¦  «        }|g d¢k    rd|D ]a}|j        d	k    rT| j                              |g d
¢g d¢¦  «        }|�2|d         j        d         |j        d         k    r|d         | _         dS Œbt          |¦  «        dk    rÄ|d         j        dk    r³|d         j        d         |v rž||d         j        d                  }t          |¦  «        dk    rr|d         j        dk    ra|d         j        d         |v rL||d         j        d                  }	|	D ]}|j        dk    r
|| _         dS Œt	          d„ |	D ¦   «         ¦  «        }|r5|g d¢k    r,|g d¢k    r$|g d¢k    rt                               d¦  «         dS n,|g d¢k    r$|g d¢k    rt                               d¦  «         dS dS )a§  Check that LayerNormalization has a child of Attention node or subgraph like Attention.

        Args:
            layernorm (NodeProto): LayerNormalization node
            input_name_to_nodes (Dict[str, List[NodeProto]]): map from input name to nodes
            is_distil_bert (bool): whether it is DistilBert or not

        Returns:
            bool: whether there is Attention node or subgraph like Attention
        Ú	AttentionF)Ú	recursiveNTr   c                 ó   — g | ]	}|j         ‘Œ
S © ©Úop_type©Ú.0Úchilds     r   ú
<listcomp>zCFusionEmbedLayerNoMask.check_attention_subgraph.<locals>.<listcomp>J   s   € Ð EÐ EÐ E°5 ¤Ð EÐ EÐ Er   )ÚMatMulr5   r5   r   r   )ÚAddr5   ÚMultiHeadAttentionr5   )NNr   r   éÿÿÿÿé   r"   r5   r6   c                 ó   — g | ]	}|j         ‘Œ
S r.   r/   r1   s     r   r4   zCFusionEmbedLayerNoMask.check_attention_subgraph.<locals>.<listcomp>g   s   € Ð(JÐ(JÐ(J¸5¨¬Ð(JÐ(JÐ(Jr   )r5   r5   r5   ÚShaper   )r6   r5   r5   r5   r;   r;   )r6   r5   r5   r5   r;   z<No Attention like subgraph in children of LayerNormalization)r6   r5   r5   r5   )r   Úfind_first_child_by_typer   ÚoutputÚsortedr0   r#   ÚinputÚcross_attentionÚlenÚloggerÚdebug)
r   r'   r(   r)   ÚchildrenÚchildren_typesÚnodeÚpath1ÚgrandchildrenÚnodess
             r   Úcheck_attention_subgraphz/FusionEmbedLayerNoMask.check_attention_subgraph0   sÂ  € ð  œ×<Ò<Ø�{Ð$7À5ð =ñ 
ô 
ˆŒð Œ>Ð%Ø�4àÔ˜AÔÐ&9Ð9Ð9Ø�5Ø& yÔ'7¸Ô':Ô;ˆÝÐ EÐ E¸HÐ EÑ EÔ EÑFÔFˆð ÐUÐUÐUÒUÐUØ ð 	$ð 	$�Ø”<Ð#;Ò;Ð;Ø œJ×8Ò8ØØIÐIÐIØ*Ð*Ð*ñô �Eð
 Ð(¨U°2¬Y¬_¸QÔ-?À9ÔCSÐTUÔCVÒ-VÐ-VØ/4°Q¬x˜Ô,Ø#˜t˜tøõ ˆx‰=Œ=˜AÒÐ (¨1¤+Ô"5¸Ò"AÐ"AÀhÈqÄkÔFXÐYZÔF[Ð_rÐFrÐFrØ/°¸´Ô0BÀ1Ô0EÔFˆMå�MÑ"Ô" aÒ'Ð'Ø! !Ô$Ô,°Ò5Ð5Ø! !Ô$Ô+¨AÔ.Ð2EÐEÐEà+¨M¸!Ô,<Ô,CÀAÔ,FÔG�Ø!ð $ð $�DØ”| {Ò2Ð2Ø)-˜œØ#˜t˜tð 3õ "(Ð(JÐ(JÀEÐ(JÑ(JÔ(JÑ!KÔ!K�ð ð 	ð Ð"cÐ"cÐ"cÒcÐcØ"Ð&]Ð&]Ð&]Ò]Ð]Ø"Ð&TÐ&TÐ&TÒTÐTå—’Ð[Ñ\Ô\Ð\Ø�uøàð "ð "ð "ò ð ð
 !ð %ð %ð %ò ð õ —’Ð[Ñ\Ô\Ð\Ø�uàˆtr   c                 óB  — | j                              |ddgddg¦  «        }|€$| j                              |g d¢g d¢¦  «        }|€dS |d         |d	         }}|j        d         |k    rdS | j                              |g d
¢g d¢fg d¢g d¢fg|¦  «        \  }}}|€dS |d         }	| j                             |	dd¦  «        r| j                             |	dd¦  «        sdS |d         }
| j                             |
dd¦  «        sdS |d	         }|j        d         |k    rdS dS )az    Match position embedding path from input_ids to Gather for DistilBert.

        Pattern is like the following:
                 (input_ids)
                      |
                     Shape
                       |                          |    Gather (indices=1)
                       |       |
                       |      Cast (optional)
                       |       |
                       |      Range (start=0, end=*, delta=1)
                       |       |
                       |    Unsqueeze
                       |    /
                      Expand
                        |
                      Gather
        ÚExpandr;   r"   N)rL   ÚWhereÚReshaper;   )r"   r"   r9   r   Fr   r8   )Ú	UnsqueezeÚRangeÚCastr!   r;   )r   r   r"   r   r   )rO   rP   r!   r;   )r   r   r"   r   r9   éþÿÿÿT)r   r#   r?   Úmatch_parent_pathsr   Úcheck_node_input_value)r   Úposition_embedding_gatherÚ	input_idsÚoutput_name_to_noderG   ÚexpandÚshapeÚ_Úpath2Ú
range_nodeÚgather_nodeÚ
shape_nodes               r   Ú#match_position_embedding_distilbertz:FusionEmbedLayerNoMask.match_position_embedding_distilbert„   s‚  € ð* ”
×,Ò,Ð-FÈÐSZÐH[Ð^_ÐabÐ]cÑdÔdˆØˆ=Ø”J×0Ò0Ø)Ø7Ð7Ð7Ø��ñô ˆEð
 ˆ}Ø�uà˜aœ %¨¤)�ˆØŒ;�qŒ>˜YÒ&Ð&Ø�5à”j×3Ò3ØàBÐBÐBÀOÀOÀOÐTØ:Ð:Ð:¸L¸L¸LÐIðð  ñ
ô 
‰ˆˆ5�!ð ˆ=Ø�5à˜1”Xˆ
àŒJ×-Ò-¨j¸!¸QÑ?Ô?ð	ØDHÄJ×DeÒDeÐfpÐrsÐuvÑDwÔDwð	ð �5à˜B”iˆØ”
×1Ò1°+¸qÀ!ÑDÔDð 	Ø�5à˜2”Yˆ
ØÔ˜AÔ )Ò+Ð+Ø�5àˆtr   c                 ó   — dS )aY  Match position embedding path from input_ids to Gather for Roberta.

        Roberta Embedding Layer Pattern (* is optional since it might be removed by ORT, ? is the padding word id):
          (input_ids) --> Equal(B=?) -- Not -- Cast(to=6) -- CumSum(axis=1) -- Mul -- Cast(to=7) -- Add(B=1) -- Cast(to=7)* --> Gather
                                                |                              ^
                                                V                              |
                                                +------------------------------+

        Roberta new pattern from transformers v4.9:
           (input_ids) --> Equal(B=?) -- Not -- Cast(to=6) -- CumSum(axis=1) -- Add(B=0) -- Mul -- Cast(to=7) -- Add(B=1) --> Gather
                                                |                                           ^
                                                V                                           |
                                                +-------------------------------------------+

        start_node = position_embedding_gather
        start_index = 1

        # match optional Cast node.
        parent = self.model.get_parent(start_node, start_index, output_name_to_node)
        if parent is None:
            return
        if parent.op_type == "Cast":
            if OnnxModel.get_node_attribute(parent, "to") != 7:
                return
            start_node = parent
            start_index = 0

        i, path, return_indices = self.model.match_parent_paths(
            start_node,
            [ (['Add', 'Cast', 'Mul', 'CumSum', 'Cast', 'Not', 'Equal'], [start_index, 0, 0, 0, 0, 0, 0]),
              (['Add', 'Cast', 'Mul', 'Add', 'CumSum', 'Cast', 'Not', 'Equal'], [start_index, 0, 0, 0, 0, 0, 0, 0])],
            output_name_to_node)

        if path is not None:
            # constant input of Add shall be 1.
            i, value = self.model.get_constant_input(path[0])
            if value != 1:
                return False

            _, self.padding_word_id = self.model.get_constant_input(path[-1])

            return input_ids == path[-1].input[0]
        Fr.   ©r   rU   rV   rW   s       r   Ú match_position_embedding_robertaz7FusionEmbedLayerNoMask.match_position_embedding_robertaÂ   s
   € ðZ ˆur   c                 ó*  — | j                              |ddgddg|¦  «        }|€dS |\  }}| j                              |j        d         ¦  «        }|�˜t	          |j        ¦  «        dk    r€|j        d         dk    ro| j                             |ddg¦  «        rR| j                             |ddg¦  «        r5t	          |j        ¦  «        d	k    s| j                             |d	dg¦  «        sdS | j                              ¦   «         }|d
k     rt          j
        |ddg¦  «        sdS n| j                             |ddg¦  «        sdS | j                              |d|¦  «        }	|	€dS |	j        dk    r;| j                             |	dd¦  «        sdS | j                              |	d|¦  «        }
n|	}
|
�|
j        dk    rdS | j                             |
dd¦  «        sdS | j                              |
d|¦  «        }|�|j        dk    rdS ||j        d         k    S )a	    Match position embedding path from input_ids to Gather for BERT.

        BERT Embedding Layer Pattern:
                                    (input_ids)
                                   /                                          /          Shape
                                /              |
                              /              Gather (indices=1)
                             /                  |
                            /                  Add (optional, B=0)
                           /                    |
                        Gather (segment_ids) Unsqueeze (axes=0)
                           \        |           |
                            \     Gather      Slice (data[1,512], starts=0, ends=*, axes=1, steps=1)
                              \    /            |
                                Add          Gather
                                   \       /
                                      Add
                                       |
                                LayerNormalization
        ÚSlicerO   r"   r9   NFr   é   é   é   Úaxesr6   r!   r;   )r   r#   Úget_constant_valuer?   rA   rY   r   rT   Úget_opset_versionr   Úcheck_node_attributeÚ
get_parentr0   )r   rU   rV   rW   ÚpathÚsliceÚ	unsqueezeÚslice_weightÚopset_versionrF   ÚgatherrY   s               r   Úmatch_position_embedding_bertz4FusionEmbedLayerNoMask.match_position_embedding_bertñ   sW  € ð, Œz×+Ò+Ø%Ø�kÐ"Ø�ˆFØñ	
ô 
ˆð ˆ<Ø�5àÑˆˆyØ”z×4Ò4°U´[À´^ÑDÔDˆàÐ$Ý�LÔ&Ñ'Ô'¨1Ò,Ð,ØÔ" 1Ô%¨Ò*Ð*Ø”
×1Ò1°%¸¸Q¸CÑ@Ô@ð +à”
×1Ò1°%¸¸Q¸CÑ@Ô@ð +õ �U”[Ñ!Ô! QÒ&Ð&¨$¬*×*KÒ*KÈEÐSTÐWXÐVYÑ*ZÔ*ZÐ&à�5àœ
×4Ò4Ñ6Ô6ˆØ˜2ÒÐÝÔ3°I¸vÈÀsÑKÔKð Ø�uðð ”:×4Ò4°YÀÀAÀ3ÑGÔGð Ø�uàŒz×$Ò$ Y°Ð3FÑGÔGˆØˆ<Ø�5ØŒ<˜5Ò Ð Ø”:×4Ò4°T¸1¸aÑ@Ô@ð Ø�uØ”Z×*Ò*¨4°Ð4GÑHÔHˆFˆFàˆFàˆ>˜Vœ^¨xÒ7Ð7Ø�5Ø”
×1Ò1°&¸!¸QÑ?Ô?ð 	Ø�5à”
×%Ò% f¨aÐ1DÑEÔEˆØˆ=˜EœM¨WÒ4Ð4Ø�5à˜EœK¨œNÒ*Ð*r   c                 ój   — |                       |||¦  «        rdS |                      |||¦  «        rdS dS )NTF)rs   r_   ra   s       r   Úmatch_position_embeddingz/FusionEmbedLayerNoMask.match_position_embedding9  sK   € Ø×-Ò-Ð.GÈÐTgÑhÔhð 	Ø�4ð ×3Ò3Ð4MÈyÐZmÑnÔnð 	Ø�4àˆur   c                 óÊ  — |j         d         }|r|j         d         nd}|j         d         }| j        s'| j                             d¬¦  «        | _        d| _        | j        �ë| j                             |¦  «        }| j                             |¦  «        }|r|sJ ‚t          |¦  «        dk    r%t          |¦  «        dk    r|d         |d         k    s"t                               d|› d|› �¦  «         dS |rU| j         	                    ||¦  «        s:t                               d	|› d
| j                             |¦  «        › �¦  «         dS | j         
                    |j         d         ¦  «        }	|	�t          |	j        ¦  «        dk    rt                               d¦  «         dS | j         
                    |j         d         ¦  «        }
|
�4t          |
j        ¦  «        dk    s|	j        d         |
j        d         k    rt                               d¦  «         dS |rw| j         
                    |j         d         ¦  «        }|�4t          |j        ¦  «        dk    s|	j        d         |j        d         k    rt                               d¦  «         dS |	j        d         |
j        d         k    rRt                               d|j         d         › d|	j        d         › d|j         d         › d|
j        d         › �¦  «         |rÜ|	j        d         |j        d         k    rRt                               d|j         d         › d|	j        d         › d|j         d         › d|j        d         › �¦  «         |
j        d         |j        d         k    rRt                               d|j         d         › d|
j        d         › d|j         d         › d|j        d         › �¦  «         dS )zXSanity check of embedding weights, and match hidden_size of weights and shape of inputs.r"   NT)Úupdater9   z^Cannot fuse EmbedLayerNormalization: input_ids and position_ids not matched in 2nd dimension: z vs FzYCannot fuse EmbedLayerNormalization: input_ids and segment_ids does not have same shape: z != r   zICannot fuse EmbedLayerNormalization: word embedding table is not expectedzMCannot fuse EmbedLayerNormalization: position embedding table is not expectedzLCannot fuse EmbedLayerNormalization: segment embedding table is not expectedzword_embedding_table (z) size z <= position_embedding_table (z <= segment_embedding_table (zposition_embedding_table ()r?   r   r   Úinfer_runtime_shaper   Úget_edge_shaperA   rB   ÚinfoÚcompare_shaperi   rY   Úwarning)r   Úword_embedding_gatherÚsegment_embedding_gatherrU   rV   Úsegment_idsÚposition_idsÚinput_ids_shapeÚposition_ids_shapeÚword_embedding_tableÚposition_embedding_tableÚsegment_embedding_tables               r   Úcheck_embeddingz&FusionEmbedLayerNoMask.check_embeddingG  s>  € à)Ô/°Ô2ˆ	Ø;SÐ]Ð.Ô4°QÔ7Ð7ÐY]ˆØ0Ô6°qÔ9ˆàÔ$ð 	)Ø#œz×=Ò=ÀTÐ=ÑJÔJˆDÔØ$(ˆDÔ!àÔÐ'Ø"Ô.×=Ò=¸iÑHÔHˆOØ!%Ô!1×!@Ò!@ÀÑ!NÔ!NÐØ"Ð9Ð'9Ð9Ð9Ð9å�OÑ$Ô$¨Ò)Ð)ÝÐ*Ñ+Ô+¨qÒ0Ð0Ø# AÔ&Ð*<¸QÔ*?Ò?Ð?å—’ð _ð  vEð  _ð  _ð  K]ð  _ð  _ñô ð ð �uàð  4Ô#3×#AÒ#AÀ)È[Ñ#YÔ#Yð Ý—’ð tÐpð  tð  tð  FJô  FV÷  Feò  Feð  fqñ  Frô  Frð  tð  tñô ð ð �uà#œz×<Ò<Ð=RÔ=XÐYZÔ=[Ñ\Ô\ÐØÐ'­3Ð/CÔ/IÑ+JÔ+JÈaÒ+OÐ+OÝ�KŠKÐcÑdÔdÐdØ�5à#'¤:×#@Ò#@ÐAZÔA`ÐabÔAcÑ#dÔ#dÐ à$Ð,ÝÐ+Ô1Ñ2Ô2°aÒ7Ð7Ø$Ô*¨1Ô-Ð1IÔ1OÐPQÔ1RÒRÐRå�KŠKÐgÑhÔhÐhØ�5àð 	Ø&*¤j×&CÒ&CÐD\ÔDbÐcdÔDeÑ&fÔ&fÐ#à'Ð/ÝÐ.Ô4Ñ5Ô5¸Ò:Ð:Ø(Ô.¨qÔ1Ð5LÔ5RÐSTÔ5UÒUÐUå—’ÐjÑkÔkÐkØ�uð  Ô% aÔ(Ð,DÔ,JÈ1Ô,MÒMÐMÝ�NŠNð \Ð)>Ô)DÀQÔ)Gð  \ð  \ÐPdÔPjÐklÔPmð  \ð  \ð  Ngô  Nmð  noô  Npð  \ð  \ð  yQô  yWð  XYô  yZð  \ð  \ñô ð ð ð 		Ø#Ô)¨!Ô,Ð0GÔ0MÈaÔ0PÒPÐPÝ—’ð ]Ð-BÔ-HÈÔ-Kð  ]ð  ]ÐThÔTnÐopÔTqð  ]ð  ]ð  Qiô  Qoð  pqô  Qrð  ]ð  ]ð  {Rô  {Xð  YZô  {[ð  ]ð  ]ñô ð ð (Ô-¨aÔ0Ð4KÔ4QÐRSÔ4TÒTÐTÝ—’ð iÐ1JÔ1PÐQRÔ1Sð  ið  iÐ\tÔ\zÐ{|Ô\}ð  ið  ið  ]uô  ]{ð  |}ô  ]~ð  ið  ið  G^ô  Gdð  efô  Ggð  ið  iñô ð ð ˆtr   Ú
input_namec                 ó   — d}| j                              |¦  «        }|�@|j        j        j        t
          j        k    r| j                             |¦  «        \  }}n |}n| j                             |¦  «        \  }}||fS )a¨  Cast a graph input or node input to int32.

        Args:
            input_name (str): name of graph input or node input

        Returns:
            A tuple of casted input name and the cast node.
            int32_output (str): If input is int32, it is the input name, Otherwise it is output name of Cast node.
            input_cast_node (Union[None, NodeProto]): Cast node. It could be None if input is int32.
        N)	r   Úfind_graph_inputÚtypeÚtensor_typeÚ	elem_typer   ÚINT32r   Úcast_input_to_int32)r   r‡   Úinput_cast_nodeÚgraph_inputÚint32_outputs        r   Úcast_to_int32z$FusionEmbedLayerNoMask.cast_to_int32‘  sƒ   € ð ˆØ”j×1Ò1°*Ñ=Ô=ˆØÐ"ØÔÔ+Ô5½Ô9JÒJÐJØ04´
×0NÒ0NÈzÑ0ZÔ0ZÑ-�˜o˜oà)��à,0¬J×,JÒ,JÈ:Ñ,VÔ,VÑ)ˆL˜/à˜_Ð,Ð,r   FrV   r}   rU   r~   r€   c	                 ót  — g }	|                       |¦  «        \  }}
| j                             d¦  «        }|j        dk    r|j        d         }|j        d         }n|j        d         }|j        d         }d}|�N|                       |j        d         ¦  «        \  }}
|||j        d         |j        d         |j        d         ||g}n|d|j        d         |j        d         d||g}|�B|                     d¦  «         |                       |¦  «        \  }}
|                     |¦  «         |d	z   |d
z   g}|r|�|n|dz   }|                     |¦  «         t          j        d|||¬¦  «        }d|_        |j	        D ](}|j
        dk    r|j	                             |g¦  «         Œ)t          |j	        ¦  «        dk    r.|j	                             t          j        dd¦  «        g¦  «         |	                     |¦  «         |	D ]}| j        | j        |j
        <   Œ| j                             |	¦  «         || _        |S )ag  Create an EmbedLayerNormalization node. Note that segment embedding is optional.

        Args:
            input_ids (str): input_ids for word embeddings
            layernorm (NodeProto): LayerNormalization or SkipLayerNormalization node.
            word_embedding_gather (NodeProto): the Gather node for word embedding
            position_embedding_gather (NodeProto): the Gather node for position embedding
            segment_embedding_gather (Union[None, NodeProto]): the Gather node for segment embedding, or None.

        Returns:
            NodeProto: the EmbedLayerNormalization node created.
        r   r   r"   r9   re   Nr   Ú Ú_outputÚ_dummy_mask_indexÚ_embedding_sum)ÚoutputsÚnamezcom.microsoftÚepsilongê-�™—q=)r’   r   Úcreate_node_namer0   r?   Úappendr   Ú	make_nodeÚdomainÚ	attributer™   ÚextendrA   Úmake_attributeÚthis_graph_nameÚnode_name_to_graph_nameÚnodes_to_addr   )r   rV   r'   r}   rU   r~   r€   Úembedding_sum_outputÚembedding_sum_namer¤   rZ   Ú	node_nameÚgammaÚbetaÚembed_node_inputsr   Úembed_node_outputsr™   r   ÚattrF   s                        r   Úcreate_fused_nodez(FusionEmbedLayerNoMask.create_fused_node¨  s¢  € ð. ˆØ×)Ò)¨)Ñ4Ô4‰ˆ	�1à”J×/Ò/Ð0IÑJÔJˆ	àÔÐ 4Ò4Ð4Ø”O AÔ&ˆEØ”? 1Ô%ˆDˆDà”O AÔ&ˆEØ”? 1Ô%ˆDà ÐØ#Ð/Ø!×/Ò/Ð0HÔ0NÈqÔ0QÑRÔR‰NˆK˜ð ØØ%Ô+¨AÔ.Ø)Ô/°Ô2Ø(Ô.¨qÔ1ØØð!ÐÐð ØØ%Ô+¨AÔ.Ø)Ô/°Ô2ØØØð!Ðð Ð#à×$Ò$ RÑ(Ô(Ð(Ø"×0Ò0°Ñ>Ô>‰OˆL˜!Ø×$Ò$ \Ñ2Ô2Ð2à'¨)Ñ3°YÐATÑ5TÐUÐØð 	,Ø);Ð)GÐ%Ð%ÈYÐYiÑMiˆDØ×%Ò% dÑ+Ô+Ð+åÔ%Ø%ØØ&Øð	
ñ 
ô 
ˆ
ð ,ˆ
Ôð Ô&ð 	3ð 	3ˆCØŒx˜9Ò$Ð$ØÔ$×+Ò+¨S¨EÑ2Ô2Ð2øõ ˆzÔ#Ñ$Ô$¨Ò)Ð)ØÔ ×'Ò'­Ô)>¸yÈ'Ñ)RÔ)RÐ(SÑTÔTÐTð 	×Ò˜JÑ'Ô'Ð'Ø ð 	Kð 	KˆDØ6:Ô6JˆDÔ(¨¬Ñ3Ð3ØÔ× Ò  Ñ.Ô.Ð.à$ˆŒØÐr   c                 óv   — | j                              |j        d         |j        d         ¦  «         d| _        d S )Nr   T)r   Úreplace_input_of_all_nodesr=   Úprune_graph)r   r'   r   s      r   Úfinish_fusionz$FusionEmbedLayerNoMask.finish_fusion
  s9   € ØŒ
×-Ò-¨iÔ.>¸qÔ.AÀ:ÔCTÐUVÔCWÑXÔXÐXàˆÔÐÐr   c                 ó„   — |j         dk    o5t          |j        ¦  «        dk    ot          |j        d         ¦  «        dk    S )Nr   re   r   )r0   rA   r=   )r   rF   s     r   Ú"is_skip_layer_norm_with_sum_outputz9FusionEmbedLayerNoMask.is_skip_layer_norm_with_sum_output  sC   € Ø”Ð 8Ò8Ðn½cÀ$Ä+Ñ>NÔ>NÐQRÒ>RÐnÕWZÐ[_Ô[fÐghÔ[iÑWjÔWjÐmnÒWnÐnr   c           
      ór  — |                       |¦  «        }|€dS |\  }}|j        d         }	|j        d         }
|                      ||d¬¦  «        sdS |                      |d |¦  «        sdS |j        dk    rK|                      |¦  «        }d}|}|r|j        d         nd }|d uo| j                             |¦  «        d u}nŠ|}|j        dk    rdnd}t          |j        ¦  «        |k    r|j        |         nd }|d uo| j                             |¦  «        d u}|o||v ot          ||         ¦  «        dk    }|d uo|j        dk    p|p|}|  
                    |	|||||
||r|nd ¬¦  «        }|r2d	|j        |<   |s&| j                             ||j        d
         ¦  «         |                      ||¦  «         dS )NFr"   ©r)   r   re   r6   r   )r¥   r¦   Ú_no_use__to_be_removed_r9   T)r&   r?   rJ   r†   r0   r³   r=   r   Úfind_graph_outputrA   r­   r¯   r±   )r   r'   Úadd_before_layernormr(   rW   Úoptional_segment_gatherÚ
two_gatherr}   rU   rV   r€   Úneed_embedding_sum_outputÚsum_output_indexÚnode_with_sum_outputÚ
sum_outputÚis_sum_graph_outputÚis_sum_used_by_multiple_nodesr   s                     r   Ú	fuse_gpt2z FusionEmbedLayerNoMask.fuse_gpt2  sn  € ð( ×*Ò*Ð+?Ñ@Ô@ˆ
ØÐØ�5à;EÑ8ÐÐ8Ø)Ô/°Ô2ˆ	Ø0Ô6°qÔ9ˆà×,Ò,¨YÐ8KÐ\aÐ,ÑbÔbð 	Ø�5à×#Ò#Ð$9¸4ÐAZÑ[Ô[ð 	Ø�5ð ÔÐ 8Ò8Ð8Ø(,×(OÒ(OÐPYÑ(ZÔ(ZÐ%Ø ÐØ#,Ð Ø0IÐS˜Ô)¨!Ô,Ð,ÈtˆJØ#-°TÐ#9Ð"uÀÄ
×@\Ò@\Ð]gÑ@hÔ@hÐptÐ@tÐÐà#7Ð Ø$8Ô$@ÀEÒ$IÐ$I˜q˜qÈqÐõ Ð+Ô2Ñ3Ô3Ð6FÒFÐFð %Ô+Ð,<Ô=Ð=àð ð
 $.°TÐ#9Ð"uÀÄ
×@\Ò@\Ð]gÑ@hÔ@hÐptÐ@tÐàÐo 
Ð.AÐ AÐoÅsÐK^Ð_iÔKjÑGkÔGkÐnoÒGoð *ð *4¸4Ð)?ð )Ø$Ô,°Ò5ÐmÐ9LÐmÐPmð &ð
 ×+Ò+ØØØ!Ø%Ø#ØØ!:Ø-@ÐJ˜z˜zÀdð ,ñ 	
ô 	
ˆ
ð %ð 	XØ<UÐ Ô'Ð(8Ñ9Ø&ð XØ”
×5Ò5°jÀ*ÔBSÐTUÔBVÑWÔWÐWà×Ò˜9 jÑ1Ô1Ð1Øˆtr   c                 óR  — |                       |¦  «        }|€dS |\  }}|j        d         }|                      ||d¬¦  «        sdS |                      |||¦  «        sdS |                      |d|¦  «        sdS |                      ||||d¦  «        }	|                      ||	¦  «         dS )aÄ  Fuse embedding layer for DistilBert
        Args:
            layernorm (NodeProto): node of LayerNormalization or SkipLayerNormalization
            add_before_layernorm (NodeProto): the Add node before LayerNormalization, or the SkipLayerNormalization itself
            input_name_to_nodes (Dict[str, List[NodeProto]]): map from input name to nodes
            output_name_to_node (Dict[str, List[NodeProto]]): map from output name to nodes
        NFr"   Trµ   )r&   r?   rJ   ru   r†   r­   r±   )
r   r'   r¸   r(   rW   rº   r}   rU   rV   r   s
             r   Úfuse_distilbertz&FusionEmbedLayerNoMask.fuse_distilbertc  sâ   € ð& ×*Ò*Ð+?Ñ@Ô@ˆ
ØÐØ�5à;EÑ8ÐÐ8Ø)Ô/°Ô2ˆ	à×,Ò,¨YÐ8KÐ\`Ð,ÑaÔað 	Ø�5à×,Ò,Ð-FÈ	ÐSfÑgÔgð 	Ø�5à×#Ò#Ð$9¸4ÐAZÑ[Ô[ð 	Ø�5à×+Ò+Ø�yÐ"7Ð9RÐTXñ
ô 
ˆ
ð 	×Ò˜9 jÑ1Ô1Ð1Øˆtr   c                 ó0  — | j                              |dgdg¦  «        }|€dS |                      |d         ¦  «        }|€dS |\  }}|j        d         }	|                      ||d¬¦  «        sdS | j                              |dgdg¦  «        }
|
€dS |
d         }|                      ||	|¦  «        s|                      ||	|¦  «        sdS |}|}|}|                      |||¦  «        sdS |                      |	||||¦  «        }|                      ||¦  «         dS )	a¾  Fuse embedding layer for Bert
        Args:
            layernorm (NodeProto): node of LayerNormalization or SkipLayerNormalization
            add_before_layernorm (NodeProto): the Add node before LayerNormalization, or the SkipLayerNormalization itself
            input_name_to_nodes (Dict[str, List[NodeProto]]): map from input name to nodes
            output_name_to_node (Dict[str, List[NodeProto]]): map from output name to nodes
        r6   r   NFr"   rµ   r!   T)	r   r#   r&   r?   rJ   ru   r†   r­   r±   )r   r'   r¸   r(   rW   Úadd_2_gatherrº   r}   r~   rV   Úposition_embedding_pathrU   Útempr   s                 r   Ú	fuse_bertz FusionEmbedLayerNoMask.fuse_bertŒ  sx  € ð ”z×3Ò3Ð4HÈ5È'ÐTUÐSVÑWÔWˆØÐØ�5à×*Ò*¨<¸¬?Ñ;Ô;ˆ
ØÐØ�5à:DÑ7ÐÐ7à)Ô/°Ô2ˆ	à×,Ò,¨YÐ8KÐ\aÐ,ÑbÔbð 	Ø�5à"&¤*×">Ò">Ð?SÐV^ÐU_ÐbcÐadÑ"eÔ"eÐØ"Ð*Ø�5à$;¸AÔ$>Ð!Ø×,Ò,Ð-FÈ	ÐSfÑgÔgð 	-Ø×0Ò0Ð1IÈ9ÐViÑjÔjð Ø�uà+ˆDØ'@Ð$Ø(,Ð%à×#Ò#Ð$9Ð;SÐUnÑoÔoð 	Ø�5à×+Ò+ØØØ!Ø%Ø$ñ
ô 
ˆ
ð 	×Ò˜9 jÑ1Ô1Ð1Øˆtr   c                 ó4  — | j                              |dgdg¦  «        }|j        dk    r|€d S |d         }d }n�| j                              |dgdg¦  «        }| j                              |dgdg¦  «        }|€|�|€d S |d         }|d         }n;|�5|€3| j                              |dgdg¦  «        }|€d S |d         }|d         }n|}d }|                      |||||¦  «        rd S |                      ||||¦  «        rd S |                      ||||¦  «        rd S d S )Nr6   r   r   r!   r"   )r   r#   r0   rÁ   rÃ   rÈ   )	r   rF   r(   rW   Úfirst_add_pathr¸   r¹   r$   r%   s	            r   ÚfusezFusionEmbedLayerNoMask.fuse¾  s‰  € Øœ×5Ò5°d¸U¸GÀaÀSÑIÔIˆØŒ<Ð/Ò/Ð/ØÐ%Ø�Ø#1°!Ô#4Ð Ø&*Ð#Ð#à œJ×8Ò8¸À¸zÈAÈ3ÑOÔOˆMØ œJ×8Ò8¸À¸zÈAÈ3ÑOÔOˆMØÐ$¨Ð)BØ!Ð)Ø�FØ'5°aÔ'8Ð$Ø*7¸Ô*:Ð'Ð'ØÐ*¨}Ð/DØ!%¤×!=Ò!=¸dÀUÀGÈaÈSÑ!QÔ!Q�Ø!Ð)Ø�FØ'5°aÔ'8Ð$Ø*7¸Ô*:Ð'Ð'à'+Ð$Ø*.Ð'à�>Š>ØÐ&Ð(;Ð=PÐRiñ
ô 
ð 	ð ˆFà×Ò Ð&:Ð<OÐQdÑeÔeð 	ØˆFà�>Š>˜$Ð 4Ð6IÐK^Ñ_Ô_ð 	ØˆFð	ð 	r   )r   )NFN)N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r	   Ústrr   r   Útupler&   ÚdictÚlistÚboolrJ   r_   rb   rs   ru   r†   r’   r­   r±   r³   rÁ   rÃ   rÈ   rË   Ú__classcell__©r   s   @r   r   r      sE  ø€ € € € € ðð ð
ð ˜ið °cð ð ð ð ð ð ð	2 Ið 	2°$¸¸yÈ)Ð?SÔ9TÑ2Tð 	2ð 	2ð 	2ð 	2ðRàðRð " # t¨I¤Ð"6Ô7ðRð ð	Rð
 
ðRð Rð Rð Rðh<ð <ð <ð|-ð -ð -ð^F+ð F+ð F+ðPð ð ðHð Hð HðT-¨ð -°°c¸4À)Ñ;KÐ6KÔ0Lð -ð -ð -ð -ð< $(Ø"Øð`ð `àð`ð ð`ð  )ð	`ð
 $-ð`ð #'¨Ñ"2ð`ð ˜D‘jð`ð `ð `ð `ðD ð  ð  ð
oð oð oð rvðOð Oð Oð Oðb'ð 'ð 'ðR0ð 0ð 0ðd"ð "ð "ð "ð "ð "ð "r   r   c                   ó6   ‡ — e Zd Zddefˆ fd„Zd„ Zˆ fd„Zˆ xZS )ÚFusionEmbedLayerNormalizationFr   c                 óZ   •— t          ¦   «                              |d¦  «         || _        d S )Nz	with mask)r   r   Úuse_mask_index)r   r   rÚ   r   s      €r   r   z&FusionEmbedLayerNormalization.__init__ä  s+   ø€ Ý‰Œ×Ò˜ Ñ,Ô,Ð,Ø,ˆÔÐÐr   c                 ój  — | j         }t          |j        ¦  «        dk    r;|j                             |¦  «         t                               d|j        ¦  «         nrt          |j        ¦  «        dk    r8|j        d         s+||j        d<   t                               d|j        ¦  «         n"t                               d|j        ¦  «         d S |D ]c}t                               d|j        ¦  «         |j        dk    r|j        d         |j        d<   ŒC|j        d	k    r|j        d         |j        d
<   Œdd S )Né   zappend mask to %szreplace mask in %szskip mask in %szupdate mask_index in %sr+   r"   re   r7   rf   )	r   rA   r?   rœ   rB   rC   r™   r0   r=   )r   Ú
mask_int32Úattention_nodesr   Úattention_nodes        r   Úreplace_maskz*FusionEmbedLayerNormalization.replace_maskè  s4  € ð ”_ˆ
ÝˆzÔÑ Ô  AÒ%Ð%ØÔ×#Ò# JÑ/Ô/Ð/Ý�LŠLÐ,¨j¬oÑ>Ô>Ð>Ð>Ý�Ô!Ñ"Ô" QÒ&Ð&¨zÔ/?ÀÔ/BÐ&Ø",ˆJÔ˜QÑÝ�LŠLÐ-¨z¬Ñ?Ô?Ð?Ð?å�LŠLÐ*¨J¬OÑ<Ô<Ð<ØˆFà-ð 	?ð 	?ˆNÝ�LŠLÐ2°NÔ4GÑHÔHÐHØÔ%¨Ò4Ð4Ø*4Ô*;¸AÔ*>�Ô$ QÑ'Ð'ØÔ'Ð+?Ò?Ð?Ø*4Ô*;¸AÔ*>�Ô$ QÑ'øð	?ð 	?r   c                 ó*  •— d | _         d | _        d | _        t          ¦   «                              |||¦  «         | j        €d S | j        s1t                               d¦  «         |                      d¦  «         d S | j         €8| j        €1t                               d¦  «         |                      d¦  «         d S | j         r| j         j	        d         }n| j        j	        d         }||         }| j
                             |¦  «        r9d„ |D ¦   «         }|                      ||¦  «         |                      d¦  «         d S ||vr2t                               d|¦  «         |                      d¦  «         d S ||         }|j        d	v r‹d
„ |D ¦   «         }|j        dk    rG|j	        d         }t          |¦  «        t          |¦  «        k    r| j                             |¦  «         |                      ||¦  «         |                      d¦  «         d S d S )NzG--use_mask_index is not set: EmbedLayerNormalization will not have maskz EmbedLayerNormalization(no mask)zLEmbedLayerNormalization will not have mask since attention node is not foundre   rf   c                 ó$   — g | ]}|j         d v ¯|‘ŒS ©)r+   r7   r/   ©r2   rF   s     r   r4   z6FusionEmbedLayerNormalization.fuse.<locals>.<listcomp>  ó%   € ÐvÐvÐv¨À$Ä,ÐRuÐBuÐBu˜tÐBuÐBuÐBur   z"EmbedLayerNormalization(with mask)zHEmbedLayerNormalization will not have mask since %s is not a node output)Ú	ReduceSumrQ   c                 ó$   — g | ]}|j         d v ¯|‘ŒS rã   r/   rä   s     r   r4   z6FusionEmbedLayerNormalization.fuse.<locals>.<listcomp>$  rå   r   ræ   r   )r   r@   r   r   rË   rÚ   rB   rC   Úincrease_counterr?   r   r‰   rà   r0   rA   Únodes_to_removerœ   )r   rF   r(   rW   rÝ   Úchildren_nodesrÞ   r   s          €r   rË   z"FusionEmbedLayerNormalization.fuseý  s1  ø€ àˆŒØ#ˆÔØˆŒÝ‰Œ�Š�TÐ.Ð0CÑDÔDÐDàŒ?Ð"ØˆFàÔ"ð 	Ý�LŠLÐbÑcÔcÐcØ×!Ò!Ð"DÑEÔEÐEØˆFàŒ>Ð! dÔ&:Ð&BÝ�LŠLÐgÑhÔhÐhØ×!Ò!Ð"DÑEÔEÐEØˆFàŒ>ð 	7ØœÔ-¨aÔ0ˆJˆJàÔ-Ô3°AÔ6ˆJà,¨ZÔ8ˆØŒ:×&Ò& zÑ2Ô2ð 	ØvÐv°ÐvÑvÔvˆOØ×Ò˜j¨/Ñ:Ô:Ð:Ø×!Ò!Ð"FÑGÔGÐGØˆFàÐ0Ð0Ð0Ý�LŠLÐcÐeoÑpÔpÐpØ×!Ò!Ð"DÑEÔEÐEØˆFà" :Ô.ˆØŒ<Ð0Ð0Ð0ØvÐv°ÐvÑvÔvˆOØŒ|˜{Ò*Ð*Ø!œZ¨œ]�
Ý�~Ñ&Ô&­#¨oÑ*>Ô*>Ò>Ð>ØÔ(×/Ò/°Ñ5Ô5Ð5Ø×Ò˜j¨/Ñ:Ô:Ð:Ø×!Ò!Ð"FÑGÔGÐGÐGÐGð 1Ð0r   )F)rÌ   rÍ   rÎ   r	   r   rà   rË   rÕ   rÖ   s   @r   rØ   rØ   ã  sz   ø€ € € € € ð-ð -˜ið -ð -ð -ð -ð -ð -ð?ð ?ð ?ð*-Hð -Hð -Hð -Hð -Hð -Hð -Hð -Hð -Hr   rØ   N)Úloggingr   Úfusion_baser   Úfusion_utilsr   Úonnxr   r   r   Ú
onnx_modelr	   rÌ   rB   r   rØ   r.   r   r   ú<module>rð      sø   ðð Ð Ð Ð Ð Ð à Ð Ð Ð Ð Ð Ø $Ð $Ð $Ð $Ð $Ð $Ø /Ð /Ð /Ð /Ð /Ð /Ð /Ð /Ð /Ð /Ø  Ð  Ð  Ð  Ð  Ð  à	ˆ�8Ñ	Ô	€ðPð Pð Pð Pð P˜Vñ Pô Pð PðfGHð GHð GHð GHð GHÐ$:ñ GHô GHð GHð GHð GHr   