Ë
    ÝÍ: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                   ó  ‡ — 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.
    ÚmodelÚdescriptionc                 ó†   •— t         ‰| �  |dddg|«       t        |«      | _        d | _        d| _        d | _        d | _        y )NÚEmbedLayerNormalizationÚLayerNormalizationÚSkipLayerNormalizationF)ÚsuperÚ__init__r   ÚutilsÚshape_inferÚshape_infer_doneÚ	attentionÚ
embed_node)Úselfr   r   Ú	__class__s      €ú/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/onnxruntime/transformers/fusion_embedlayer.pyr   zFusionEmbedLayerNoMask.__init__   sP   ø€ Ü‰ÑØØ%Ø!Ð#;Ð<Øô		
ô ! Ó'ˆŒ
ØˆÔØ %ˆÔð ˆŒØˆ�ó    ÚaddÚreturnNc                 óž   — | j                   j                  |dgdg«      }|€y | j                   j                  |dgdg«      }|€y |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%   sa   € ØŸ
™
×4Ñ4°S¸8¸*ÀqÀcÓJˆØÐ ØàŸ
™
×4Ñ4°S¸8¸*ÀqÀcÓJˆØÐ Øà˜QÑ ¨qÑ!1Ð1Ð1r   Ú	layernormÚinput_name_to_nodesÚis_distil_bertc                 ó4  — | j                   j                  |d|d¬«      | _        | j                  �y|j                  d   |vry||j                  d      }t	        |D �cg c]  }|j
                  ‘Œ c}«      }|g d¢k(  ri|D ]d  }|j
                  dk(  sŒ| j                   j                  |g d¢g d	¢«      }|€Œ7|d
   j                  d   |j                  d   k(  sŒZ|d   | _         y t        |«      dk(  rÆ|d   j
                  dk(  r´|d   j                  d   |v r ||d   j                  d      }	t        |	«      dk(  r}|	d   j
                  dk(  rk|	d   j                  d   |v rW||	d   j                  d      }
|
D ]  }|j
                  dk(  sŒ|| _         y t	        |
D �cg c]  }|j
                  ‘Œ c}«      }|r,|g d¢k7  r$|g d¢k7  r|g d¢k7  rt        j                  d«       yy|g d¢k7  r|g d¢k7  rt        j                  d«       yyc c}w c c}w )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)Ú	recursiveTr   )ÚMatMulr,   r,   r   r   )ÚAddr,   ÚMultiHeadAttentionr,   )NNr   r   éÿÿÿÿé   r!   r,   r-   )r,   r,   r,   ÚShaper   )r-   r,   r,   r,   r1   r1   )r-   r,   r,   r,   r1   z<No Attention like subgraph in children of LayerNormalization)r-   r,   r,   r,   )r   Úfind_first_child_by_typer   ÚoutputÚsortedÚop_typer"   ÚinputÚcross_attentionÚlenÚloggerÚdebug)r   r&   r'   r(   ÚchildrenÚchildÚchildren_typesÚnodeÚpath1ÚgrandchildrenÚnodess              r   Úcheck_attention_subgraphz/FusionEmbedLayerNoMask.check_attention_subgraph0   sR  € ð  Ÿ™×<Ñ<Ø�{Ð$7À5ð =ó 
ˆŒð �>‰>Ð%Øà×Ñ˜AÑÐ&9Ñ9ØØ& y×'7Ñ'7¸Ñ':Ñ;ˆÜ¸HÖ E°5 §£Ò EÓFˆð ÒUÒUØ ò 	$�Ø—<‘<Ð#;Ó;Ø ŸJ™J×8Ñ8ØÚIÚ*ó�Eð
 Ñ(¨U°2©Y¯_©_¸QÑ-?À9×CSÑCSÐTUÑCVÓ-VØ/4°Q©x˜Ô,Ù#ð	$ô ˆx‹=˜AÒ (¨1¡+×"5Ñ"5¸Ò"AÀhÈqÁk×FXÑFXÐYZÑF[Ð_rÑFrØ/°¸±×0BÑ0BÀ1Ñ0EÑFˆMä�MÓ" aÒ'Ø! !Ñ$×,Ñ,°Ò5Ø! !Ñ$×+Ñ+¨AÑ.Ð2EÑEà+¨M¸!Ñ,<×,CÑ,CÀAÑ,FÑG�Ø!ò $�DØ—|‘| {Ó2Ø)-˜œÙ#ð$ô "(ÀEÖ(J¸5¨¯«Ò(JÓ!K�ñ ð Ò"cÒcØ"Ò&]Ò]Ø"Ò&TÒTä—‘Ð[Ô\Øð  ð ò "ò ð
 !ò %ò ô —‘Ð[Ô\Øàùòq !Fùò: )Ks   ÁHÆ#Hc                 ó&  — | j                   j                  |ddgddg«      }|€$| j                   j                  |g d¢g d¢«      }|€y|d   |d   }}|j                  d   |k7  ry| j                   j                  |g d	¢g d
¢fg d¢g d¢fg|«      \  }}}|€y|d   }	| j                  j                  |	dd«      r| j                  j                  |	dd«      sy|d   }
| j                  j                  |
dd«      sy|d   }|j                  d   |k7  ryy)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
        ÚExpandr1   r!   )rD   ÚWhereÚReshaper1   )r!   r!   r0   r   Fr   r/   )Ú	UnsqueezeÚRangeÚCastr    r1   )r   r   r!   r   r   )rG   rH   r    r1   )r   r   r!   r   r0   éþÿÿÿT)r   r"   r6   Úmatch_parent_pathsr   Úcheck_node_input_value)r   Úposition_embedding_gatherÚ	input_idsÚoutput_name_to_noder?   ÚexpandÚshapeÚ_Úpath2Ú
range_nodeÚgather_nodeÚ
shape_nodes               r   Ú#match_position_embedding_distilbertz:FusionEmbedLayerNoMask.match_position_embedding_distilbert„   s9  € ð* —
‘
×,Ñ,Ð-FÈÐSZÐH[Ð^_ÐabÐ]cÓdˆØˆ=Ø—J‘J×0Ñ0Ø)Ú7ÚóˆEð
 ˆ}Øà˜a™ %¨¡)�ˆØ�;‰;�q‰>˜YÒ&Øà—j‘j×3Ñ3ØâBÂOÐTÚ:ºLÐIðð  ó
‰ˆˆ5�!ð ˆ=Øà˜1‘Xˆ
à�J‰J×-Ñ-¨j¸!¸QÔ?ÀDÇJÁJ×DeÑDeÐfpÐrsÐuvÔDwàà˜B‘iˆØ—
‘
×1Ñ1°+¸qÀ!ÔDØà˜2‘Yˆ
Ø×Ñ˜AÑ )Ò+Øàr   c                  ó   — y)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]
        F© ©r   rM   rN   rO   s       r   Ú match_position_embedding_robertaz7FusionEmbedLayerNoMask.match_position_embedding_robertaÂ   s   € ðZ r   c                 óN  — | j                   j                  |ddgddg|«      }|€y|\  }}| j                   j                  |j                  d   «      }|�œt	        |j
                  «      dk(  r„|j
                  d   dk(  rr| j                  j                  |ddg«      rT| j                  j                  |ddg«      r6t	        |j                  «      dk(  s| j                  j                  |ddg«      sy| j                   j                  «       }|d	k  rt        j                  |d
dg«      s y| j                  j                  |ddg«      sy| j                   j                  |d|«      }	|	€y|	j                  dk(  r<| j                  j                  |	dd«      sy| j                   j                  |	d|«      }
n|	}
|
�|
j                  dk7  ry| j                  j                  |
dd«      sy| j                   j                  |
d|«      }|�|j                  dk7  ry||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
        ÚSlicerG   r!   r0   Fr   é   é   é   Úaxesr-   r    r1   )r   r"   Úget_constant_valuer6   r8   rQ   r   rL   Úget_opset_versionr   Úcheck_node_attributeÚ
get_parentr5   )r   rM   rN   rO   ÚpathÚsliceÚ	unsqueezeÚslice_weightÚopset_versionr>   ÚgatherrQ   s               r   Úmatch_position_embedding_bertz4FusionEmbedLayerNoMask.match_position_embedding_bertñ   s  € ð, �z‰z×+Ñ+Ø%Ø�kÐ"Ø�ˆFØó	
ˆð ˆ<ØàÑˆˆyØ—z‘z×4Ñ4°U·[±[À±^ÓDˆàÐ$Ü�L×&Ñ&Ó'¨1Ò,Ø×"Ñ" 1Ñ%¨Ò*Ø—
‘
×1Ñ1°%¸¸Q¸CÔ@Ø—
‘
×1Ñ1°%¸¸Q¸CÔ@Ü�U—[‘[Ó! QÒ&¨$¯*©*×*KÑ*KÈEÐSTÐWXÐVYÔ*ZààŸ
™
×4Ñ4Ó6ˆØ˜2ÒÜ×3Ñ3°I¸vÈÀsÔKØà—:‘:×4Ñ4°YÀÀAÀ3ÔGØà�z‰z×$Ñ$ Y°Ð3FÓGˆØˆ<ØØ�<‰<˜5Ò Ø—:‘:×4Ñ4°T¸1¸aÔ@ØØ—Z‘Z×*Ñ*¨4°Ð4GÓH‰FàˆFàˆ>˜VŸ^™^¨xÒ7ØØ—
‘
×1Ñ1°&¸!¸QÔ?Øà—
‘
×%Ñ% f¨aÐ1DÓEˆØˆ=˜EŸM™M¨WÒ4Øà˜EŸK™K¨™NÑ*Ð*r   c                 óT   — | j                  |||«      ry| j                  |||«      ryy)NTF)rl   rW   rZ   s       r   Úmatch_position_embeddingz/FusionEmbedLayerNoMask.match_position_embedding9  s5   € Ø×-Ñ-Ð.GÈÐTgÔhØð ×3Ñ3Ð4MÈyÐZmÔnØàr   c                 óÖ  — |j                   d   }|r|j                   d   nd}|j                   d   }| j                  s(| j                  j                  d¬«      | _        d| _        | j                  �Ò| j                  j                  |«      }| j                  j                  |«      }|r|sJ ‚t        |«      dk(  rt        |«      dk(  r|d   |d   k(  st        j                  d|› d|› �«       y|rQ| j                  j                  ||«      s5t        j                  d	|› d
| j                  j                  |«      › �«       y| j                  j                  |j                   d   «      }	|	�t        |	j                  «      dk7  rt        j                  d«       y| j                  j                  |j                   d   «      }
|
�7t        |
j                  «      dk7  s|	j                  d   |
j                  d   k7  rt        j                  d«       y|rw| j                  j                  |j                   d   «      }|�7t        |j                  «      dk7  s|	j                  d   |j                  d   k7  rt        j                  d«       y|	j                  d   |
j                  d   k  rUt        j                  d|j                   d   › d|	j                  d   › d|j                   d   › d|
j                  d   › �«       |rè|	j                  d   j                  d   k  rUt        j                  d|j                   d   › d|	j                  d   › d|j                   d   › d|j                  d   › �«       |
j                  d   |j                  d   k  rUt        j                  d|j                   d   › d|
j                  d   › d|j                   d   › d|j                  d   › �«       y)zXSanity check of embedding weights, and match hidden_size of weights and shape of inputs.r!   NT)Úupdater0   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 ()r6   r   r   Úinfer_runtime_shaper   Úget_edge_shaper8   r9   ÚinfoÚcompare_shaperb   rQ   Úwarning)r   Úword_embedding_gatherÚsegment_embedding_gatherrM   rN   Ú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Ñ4°QÒ7ÐY]ˆØ0×6Ñ6°qÑ9ˆà×$Ò$Ø#Ÿz™z×=Ñ=ÀTÐ=ÓJˆDÔØ$(ˆDÔ!à×ÑÐ'Ø"×.Ñ.×=Ñ=¸iÓHˆOØ!%×!1Ñ!1×!@Ñ!@ÀÓ!NÐÙ"Ñ'9Ð9Ð9ä�OÓ$¨Ò)ÜÐ*Ó+¨qÒ0Ø# AÑ&Ð*<¸QÑ*?Ò?ä—‘Øtð  vEð  uFð  FJð  K]ð  J^ð  _ôð á 4×#3Ñ#3×#AÑ#AÀ)È[Ô#YÜ—‘ØoÐpð  pAð  AEð  FJ÷  FVñ  FV÷  Feñ  Feð  fqó  Frð  Esð  tôð à#Ÿz™z×<Ñ<Ð=R×=XÑ=XÐYZÑ=[Ó\ÐØÐ'¬3Ð/C×/IÑ/IÓ+JÈaÒ+OÜ�K‰KÐcÔdØà#'§:¡:×#@Ñ#@ÐAZ×A`ÑA`ÐabÑAcÓ#dÐ à$Ð,ÜÐ+×1Ñ1Ó2°aÒ7Ø$×*Ñ*¨1Ñ-Ð1I×1OÑ1OÐPQÑ1RÒRä�K‰KÐgÔhØáØ&*§j¡j×&CÑ&CÐD\×DbÑDbÐcdÑDeÓ&fÐ#à'Ð/ÜÐ.×4Ñ4Ó5¸Ò:Ø(×.Ñ.¨qÑ1Ð5L×5RÑ5RÐSTÑ5UÒUä—‘ÐjÔkØð  ×%Ñ% aÑ(Ð,D×,JÑ,JÈ1Ñ,MÒMÜ�N‰NØ(Ð)>×)DÑ)DÀQÑ)GÐ(HÈÐPd×PjÑPjÐklÑPmÐOnð  oMð  Ng÷  Nmñ  Nmð  noñ  Npð  Mqð  qxð  yQ÷  yWñ  yWð  XYñ  yZð  x[ð  \ôñ Ø#×)Ñ)¨!Ñ,Ð0G×0MÑ0MÈaÑ0PÒPÜ—‘Ø,Ð-B×-HÑ-HÈÑ-KÐ,LÈGÐTh×TnÑTnÐopÑTqÐSrð  sPð  Qi÷  Qoñ  Qoð  pqñ  Qrð  Psð  szð  {R÷  {Xñ  {Xð  YZñ  {[ð  z\ð  ]ôð (×-Ñ-¨aÑ0Ð4K×4QÑ4QÐRSÑ4TÒTÜ—‘Ø0Ð1J×1PÑ1PÐQRÑ1SÐ0TÐT[Ð\t×\zÑ\zÐ{|Ñ\}Ð[~ð  \ð  ]u÷  ]{ñ  ]{ð  |}ñ  ]~ð  \ð  Fð  G^÷  Gdñ  Gdð  efñ  Ggð  Fhð  iôð r   Ú
input_namec                 ó6  — d}| j                   j                  |«      }|�Y|j                  j                  j                  t
        j                  k7  r"| j                  j                  |«      \  }}||fS |}||fS | j                  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‘j×1Ñ1°*Ó=ˆØÐ"Ø×Ñ×+Ñ+×5Ñ5¼×9JÑ9JÒJØ04·
±
×0NÑ0NÈzÓ0ZÑ-�˜oð ˜_Ð,Ð,ð	  *�ð ˜_Ð,Ð,ð -1¯J©J×,JÑ,JÈ:Ó,VÑ)ˆL˜/à˜_Ð,Ð,r   rN   rv   rM   rw   ry   c	                 ó²  — g }	| j                  |«      \  }}
| j                  j                  d«      }|j                  dk(  r|j                  d   }|j                  d   }n|j                  d   }|j                  d   }d}|�R| j                  |j                  d   «      \  }}
|||j                  d   |j                  d   |j                  d   ||g}n#|d|j                  d   |j                  d   d||g}|�6|j                  d«       | j                  |«      \  }}
|j                  |«       |d	z   |d
z   g}|r|�|n|dz   }|j                  |«       t        j                  d|||¬«      }d|_        |j                  D ].  }|j                  dk(  sŒ|j                  j                  |g«       Œ0 t        |j                  «      dk(  r0|j                  j                  t        j                  dd«      g«       |	j                  |«       |	D ]%  }| j                  | 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!   r0   r^   Nr   Ú Ú_outputÚ_dummy_mask_indexÚ_embedding_sum)ÚoutputsÚnamezcom.microsoftÚepsilongê-�™—q=)r‹   r   Úcreate_node_namer5   r6   Úappendr   Ú	make_nodeÚdomainÚ	attributer’   Úextendr8   Úmake_attributeÚthis_graph_nameÚnode_name_to_graph_nameÚnodes_to_addr   )r   rN   r&   rv   rM   rw   ry   Úembedding_sum_outputÚembedding_sum_namer�   rR   Ú	node_nameÚgammaÚbetaÚembed_node_inputsrx   Úembed_node_outputsr’   r   Úattr>   s                        r   Úcreate_fused_nodez(FusionEmbedLayerNoMask.create_fused_node¨  s€  € ð. ˆØ×)Ñ)¨)Ó4‰ˆ	�1à—J‘J×/Ñ/Ð0IÓJˆ	à×ÑÐ 4Ò4Ø—O‘O AÑ&ˆEØ—?‘? 1Ñ%‰Dà—O‘O AÑ&ˆEØ—?‘? 1Ñ%ˆDà ÐØ#Ð/Ø!×/Ñ/Ð0H×0NÑ0NÈqÑ0QÓR‰NˆK˜ð ØØ%×+Ñ+¨AÑ.Ø)×/Ñ/°Ñ2Ø(×.Ñ.¨qÑ1ØØð!Ñð ØØ%×+Ñ+¨AÑ.Ø)×/Ñ/°Ñ2ØØØð!Ðð Ð#à×$Ñ$ RÔ(Ø"×0Ñ0°Ó>‰OˆL˜!Ø×$Ñ$ \Ô2à'¨)Ñ3°YÐATÑ5TÐUÐÙØ);Ð)GÑ%ÈYÐYiÑMiˆDØ×%Ñ% dÔ+ä×%Ñ%Ø%ØØ&Øô	
ˆ
ð ,ˆ
Ôð ×&Ñ&ò 	3ˆCØ�x‰x˜9Ó$Ø×$Ñ$×+Ñ+¨S¨EÕ2ð	3ô ˆz×#Ñ#Ó$¨Ò)Ø× Ñ ×'Ñ'¬×)>Ñ)>¸yÈ'Ó)RÐ(SÔTð 	×Ñ˜JÔ'Ø ò 	KˆDØ6:×6JÑ6JˆD×(Ñ(¨¯©Ò3ð	Kà×Ñ× Ñ  Ô.à$ˆŒØÐr   c                 ó~   — | j                   j                  |j                  d   |j                  d   «       d| _        y )Nr   T)r   Úreplace_input_of_all_nodesr3   Úprune_graph)r   r&   r   s      r   Úfinish_fusionz$FusionEmbedLayerNoMask.finish_fusion
  s5   € Ø�
‰
×-Ñ-¨i×.>Ñ.>¸qÑ.AÀ:×CTÑCTÐUVÑCWÔXàˆÕr   c                 óŽ   — |j                   dk(  xr5 t        |j                  «      dkD  xr t        |j                  d   «      dkD  S )Nr   r^   r   )r5   r8   r3   )r   r>   s     r   Ú"is_skip_layer_norm_with_sum_outputz9FusionEmbedLayerNoMask.is_skip_layer_norm_with_sum_output  sD   € Ø—‘Ð 8Ñ8Òn¼cÀ$Ç+Á+Ó>NÐQRÑ>RÒnÔWZÐ[_×[fÑ[fÐghÑ[iÓWjÐmnÑWnÐnr   c           
      ó‚  — | j                  |«      }|€y|\  }}|j                  d   }	|j                  d   }
| j                  ||d¬«      sy| j                  |d |«      sy|j                  dk(  rL| j                  |«      }d}|}|r|j                  d   nd }|d uxr | j                  j                  |«      d u}n™|}|j                  dk(  rdnd}t        |j                  «      |kD  r|j                  |   nd }|d uxr | j                  j                  |«      d u}|xr ||v xr t        ||   «      dkD  }|d uxr |j                  dk7  xs |xs |}| j                  |	|||||
||r|nd ¬«      }|r:d	|j                  |<   |s)| j                  j                  ||j                  d
   «       | j                  ||«       y)NFr!   ©r(   r   r^   r-   r   )rž   rŸ   Ú_no_use__to_be_removed_r0   T)r%   r6   rB   r   r5   r¬   r3   r   Úfind_graph_outputr8   r¦   r¨   rª   )r   r&   Úadd_before_layernormr'   rO   Úoptional_segment_gatherÚ
two_gatherrv   rM   rN   ry   Ú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  s/  € ð( ×*Ñ*Ð+?Ó@ˆ
ØÐØà;EÑ8ÐÐ8Ø)×/Ñ/°Ñ2ˆ	Ø0×6Ñ6°qÑ9ˆà×,Ñ,¨YÐ8KÐ\aÐ,ÔbØà×#Ñ#Ð$9¸4ÐAZÔ[Øð ×ÑÐ 8Ò8Ø(,×(OÑ(OÐPYÓ(ZÐ%Ø ÐØ#,Ð Ù0I˜×)Ñ)¨!Ò,ÈtˆJØ#-°TÐ#9Ò"uÀÇ
Á
×@\Ñ@\Ð]gÓ@hÐptÐ@tÑà#7Ð Ø$8×$@Ñ$@ÀEÒ$I™qÈqÐô Ð+×2Ñ2Ó3Ð6FÒFð %×+Ñ+Ð,<Ò=àð ð
 $.°TÐ#9Ò"uÀÇ
Á
×@\Ñ@\Ð]gÓ@hÐptÐ@tÐàÒo 
Ð.AÐ AÒoÄsÐK^Ð_iÑKjÓGkÐnoÑGoð *ð *4¸4Ð)?ò )Ø$×,Ñ,°Ñ5ÒmÐ9LÒmÐPmð &ð
 ×+Ñ+ØØØ!Ø%Ø#ØØ!:Ù-@™zÀdð ,ó 	
ˆ
ñ %Ø<UÐ ×'Ñ'Ð(8Ñ9Ù&Ø—
‘
×5Ñ5°jÀ*×BSÑBSÐTUÑBVÔWà×Ñ˜9 jÔ1Ør   c                 ó  — | j                  |«      }|€y|\  }}|j                  d   }| j                  ||d¬«      sy| j                  |||«      sy| j	                  |d|«      sy| j                  ||||d«      }	| j                  ||	«       y)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%   r6   rB   rn   r   r¦   rª   )
r   r&   r±   r'   rO   r³   rv   rM   rN   r   s
             r   Úfuse_distilbertz&FusionEmbedLayerNoMask.fuse_distilbertc  sµ   € ð& ×*Ñ*Ð+?Ó@ˆ
ØÐØà;EÑ8ÐÐ8Ø)×/Ñ/°Ñ2ˆ	à×,Ñ,¨YÐ8KÐ\`Ð,ÔaØà×,Ñ,Ð-FÈ	ÐSfÔgØà×#Ñ#Ð$9¸4ÐAZÔ[Øà×+Ñ+Ø�yÐ"7Ð9RÐTXó
ˆ
ð 	×Ñ˜9 jÔ1Ør   c                 óæ  — | j                   j                  |dgdg«      }|€y| j                  |d   «      }|€y|\  }}|j                  d   }	| j	                  ||d¬«      sy| j                   j                  |dgdg«      }
|
€y|
d   }| j                  ||	|«      s| j                  ||	|«      sy|}|}|}| j                  |||«      sy| j                  |	||||«      }| j                  ||«       y)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
        r-   r   Fr!   r®   r    T)	r   r"   r%   r6   rB   rn   r   r¦   rª   )r   r&   r±   r'   rO   Úadd_2_gatherr³   rv   rw   rN   Úposition_embedding_pathrM   Útempr   s                 r   Ú	fuse_bertz FusionEmbedLayerNoMask.fuse_bertŒ  s=  € ð —z‘z×3Ñ3Ð4HÈ5È'ÐTUÐSVÓWˆØÐØà×*Ñ*¨<¸©?Ó;ˆ
ØÐØà:DÑ7ÐÐ7à)×/Ñ/°Ñ2ˆ	à×,Ñ,¨YÐ8KÐ\aÐ,ÔbØà"&§*¡*×">Ñ">Ð?SÐV^ÐU_ÐbcÐadÓ"eÐØ"Ð*Øà$;¸AÑ$>Ð!Ø×,Ñ,Ð-FÈ	ÐSfÔgØ×0Ñ0Ð1IÈ9ÐViÔjØà+ˆDØ'@Ð$Ø(,Ð%à×#Ñ#Ð$9Ð;SÐUnÔoØà×+Ñ+ØØØ!Ø%Ø$ó
ˆ
ð 	×Ñ˜9 jÔ1Ør   c                 ó   — | j                   j                  |dgdg«      }|j                  dk(  r|€y |d   }d }n…| j                   j                  |dgdg«      }| j                   j                  |dgdg«      }|€|�|€y |d   }|d   }n5|�/|€-| j                   j                  |dgdg«      }|€y |d   }|d   }n|}d }| j                  |||||«      ry | j	                  ||||«      ry | j                  ||||«      ry y )Nr-   r   r   r    r!   )r   r"   r5   rº   r¼   rÁ   )	r   r>   r'   rO   Úfirst_add_pathr±   r²   r#   r$   s	            r   ÚfusezFusionEmbedLayerNoMask.fuse¾  sQ  € ØŸ™×5Ñ5°d¸U¸GÀaÀSÓIˆØ�<‰<Ð/Ò/ØÐ%ØØ#1°!Ñ#4Ð Ø&*Ñ#à ŸJ™J×8Ñ8¸À¸zÈAÈ3ÓOˆMØ ŸJ™J×8Ñ8¸À¸zÈAÈ3ÓOˆMØÐ$¨Ð)BØ!Ð)ØØ'5°aÑ'8Ð$Ø*7¸Ñ*:Ñ'ØÐ*¨}Ð/DØ!%§¡×!=Ñ!=¸dÀUÀGÈaÈSÓ!Q�Ø!Ð)ØØ'5°aÑ'8Ð$Ø*7¸Ñ*:Ñ'à'+Ð$Ø*.Ð'à�>‰>ØÐ&Ð(;Ð=PÐRiô
ð à×Ñ Ð&:Ð<OÐQdÔeØà�>‰>˜$Ð 4Ð6IÐK^Ô_Øð `r   )zno mask)NFN)N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r	   Ústrr   r   Útupler%   ÚdictÚlistÚboolrB   rW   r[   rl   rn   r   r‹   r¦   rª   r¬   rº   r¼   rÁ   rÄ   Ú__classcell__©r   s   @r   r   r      sB  ø„ ññ
˜ið °cõ ð	2 Ið 	2°$¸¸yÈ)Ð?SÑ9TÑ2Tó 	2ðRàðRð " # t¨I¡Ð"6Ñ7ðRð ð	Rð
 
óRòh<ò|-ò^F+òPòHðT-¨ð -°°c¸4À)Ñ;KÐ6KÑ0Ló -ð< $(Ø"Øñ`àð`ð ð`ð  )ð	`ð
 $-ð`ð #'¨Ñ"2ð`ð ˜D‘jó`òD ò
oð rvóOòb'òR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 )ÚFusionEmbedLayerNormalizationr   c                 ó4   •— t         ‰| �  |d«       || _        y )Nz	with mask)r   r   Úuse_mask_index)r   r   rÓ   r   s      €r   r   z&FusionEmbedLayerNormalization.__init__ä  s   ø€ Ü‰Ñ˜ Ô,Ø,ˆÕr   c                 ó²  — | j                   }t        |j                  «      dk(  r<|j                  j                  |«       t        j                  d|j                  «       nxt        |j                  «      dkD  r?|j                  d   s0||j                  d<   t        j                  d|j                  «       n!t        j                  d|j                  «       y |D ]z  }t        j                  d|j                  «       |j                  dk(  r|j                  d   |j                  d<   ŒO|j                  d	k(  sŒ_|j                  d   |j                  d
<   Œ| y )Né   zappend mask to %szreplace mask in %szskip mask in %szupdate mask_index in %sr*   r!   r^   r.   r_   )	r   r8   r6   r•   r9   r:   r’   r5   r3   )r   Ú
mask_int32Úattention_nodesr   Úattention_nodes        r   Úreplace_maskz*FusionEmbedLayerNormalization.replace_maskè  s  € ð —_‘_ˆ
Üˆz×ÑÓ  AÒ%Ø×Ñ×#Ñ# JÔ/Ü�L‰LÐ,¨j¯o©oÕ>Ü�×!Ñ!Ó" QÒ&¨z×/?Ñ/?ÀÒ/BØ",ˆJ×Ñ˜QÑÜ�L‰LÐ-¨z¯©Õ?ä�L‰LÐ*¨J¯O©OÔ<Øà-ò 	?ˆNÜ�L‰LÐ2°N×4GÑ4GÔHØ×%Ñ%¨Ò4Ø*4×*;Ñ*;¸AÑ*>�×$Ñ$ QÒ'Ø×'Ñ'Ð+?Ó?Ø*4×*;Ñ*;¸AÑ*>�×$Ñ$ QÒ'ñ	?r   c                 óH  •— d | _         d | _        d | _        t        ‰| �  |||«       | j                  €y | j
                  s't        j                  d«       | j                  d«       y | j                   €3| j                  €'t        j                  d«       | j                  d«       y | j                   r| j                   j                  d   }n| j                  j                  d   }||   }| j                  j                  |«      rB|D �cg c]  }|j                  dv sŒ|‘Œ }}| j                  ||«       | j                  d«       y ||vr(t        j                  d|«       | j                  d«       y ||   }|j                  d	v r’|D �cg c]  }|j                  dv sŒ|‘Œ }}j                  d
k(  rA|j                  d   }t        |«      t        |«      k(  r| j                  j!                  |«       | j                  ||«       | j                  d«       y y c c}w c c}w )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 foundr^   r_   )r*   r.   z"EmbedLayerNormalization(with mask)zHEmbedLayerNormalization will not have mask since %s is not a node output)Ú	ReduceSumrI   rÛ   r   )r   r7   r   r   rÄ   rÓ   r9   r:   Úincrease_counterr6   r   r‚   r5   rÙ   r8   Únodes_to_remover•   )r   r>   r'   rO   rÖ   Úchildren_nodesr×   r   s          €r   rÄ   z"FusionEmbedLayerNormalization.fuseý  sò  ø€ àˆŒØ#ˆÔØˆŒÜ‰‰�TÐ.Ð0CÔDà�?‰?Ð"Øà×"Ò"Ü�L‰LÐbÔcØ×!Ñ!Ð"DÔEØà�>‰>Ð! d×&:Ñ&:Ð&BÜ�L‰LÐgÔhØ×!Ñ!Ð"DÔEØà�>Š>ØŸ™×-Ñ-¨aÑ0‰Jà×-Ñ-×3Ñ3°AÑ6ˆJà,¨ZÑ8ˆØ�:‰:×&Ñ& zÔ2Ø0>Öv¨À$Ç,Á,ÐRuÒBuštÐvˆOÐvØ×Ñ˜j¨/Ô:Ø×!Ñ!Ð"FÔGØàÐ0Ñ0Ü�L‰LÐcÐeoÔpØ×!Ñ!Ð"DÔEØà" :Ñ.ˆØ�<‰<Ð0Ñ0Ø0>Öv¨À$Ç,Á,ÐRuÒBuštÐvˆOÐvØ�|‰|˜{Ò*Ø!ŸZ™Z¨™]�
Ü�~Ó&¬#¨oÓ*>Ò>Ø×(Ñ(×/Ñ/°Ô5Ø×Ñ˜j¨/Ô:Ø×!Ñ!Ð"FÕGð 1ùò wùò ws   Ä
HÄHÆHÆH)F)rÅ   rÆ   rÇ   r	   r   rÙ   rÄ   rÎ   rÏ   s   @r   rÑ   rÑ   ã  s    ø„ ñ-˜iõ -ò?÷*-Hð -Hr   rÑ   N)Úloggingr   Úfusion_baser   Úfusion_utilsr   Úonnxr   r   r   Ú
onnx_modelr	   rÅ   r9   r   rÑ   rY   r   r   ú<module>rä      sC   ðõ å Ý $ß /Ñ /Ý  á	�8Ó	€ôP˜Vô PôfGHÐ$:õ GHr   