Ë
    ÝÍ:jË2  ã                   ó„   — d dl mZ d dlZd dlmZmZ d dlmZmZmZm	Z	 d dl
mZ  ee«      Z G d„ d«      Z G d„ d	«      Zy)
é    )Ú	getLoggerN)Úarray_equalÚndarray)Ú	NodeProtoÚTensorProtoÚhelperÚnumpy_helper)Ú	OnnxModelc            
       ó0  — e Zd Zdefd„Zdedeeef   fd„Zddefd„Z		 	 	 ddede
d	edz  d
edz  fd„Zdefd„Zdefd„Zed„ «       Zeddefd„«       Zdededz  fd„Zed defd„«       Zedefd„«       Zed!dedefd„«       Zde
fd„Zd„ Zd„ Zd„ Zd„ Zy)"ÚFusionUtilsÚmodelc                 ó   — || _         y ©N)r   )Úselfr   s     úz/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/onnxruntime/transformers/fusion_utils.pyÚ__init__zFusionUtils.__init__   s	   € Ø %ˆ�
ó    Ú
input_nameÚreturnc                 óB  — | j                   j                  |«      }|�b|j                  j                  j                  t
        j                  k7  r1| j                  |«      \  }}t        j                  d|› d�«       d|fS t        j                  d|› d|d u› �«       d|fS )NzCasted graph input z	 to int32TzDid not cast graph input z to int32: found F)
r   Úfind_graph_inputÚtypeÚtensor_typeÚ	elem_typer   ÚINT32Úcast_input_to_int32ÚloggerÚdebug)r   r   Úgraph_inputÚcast_outputÚ	cast_nodes        r   Úcast_graph_input_to_int32z%FusionUtils.cast_graph_input_to_int32   sŸ   € Ø—j‘j×1Ñ1°*Ó=ˆØÐ" {×'7Ñ'7×'CÑ'C×'MÑ'MÔQ\×QbÑQbÒ'bØ%)×%=Ñ%=¸jÓ%IÑ"ˆK˜Ü�L‰LÐ.¨z¨l¸)ÐDÔEØ˜Ð$Ð$ä�‰Ð0°°Ð<MÈkÐaeÐNeÐMfÐgÔhØ�jÐ Ð r   c                 ó  — |dz   |z   }|dk(  rt        t        j                  «      }nI|dk(  rt        t        j                  «      }n*|dk(  rt        t        j                  «      }nt        d«      ‚| j                  |||«      }||fS )NÚ_Úint32Úfloat32Úfloat16z"Invalid target_type: {target_type})Úintr   r   ÚFLOATÚFLOAT16Ú
ValueErrorÚadd_cast_node)r   r   Útarget_typeÚoutput_nameÚto_typer!   s         r   Ú
cast_inputzFusionUtils.cast_input   s„   € Ø  3Ñ&¨Ñ4ˆà˜'Ò!Üœ+×+Ñ+Ó,‰GØ˜IÒ%Üœ+×+Ñ+Ó,‰GØ˜IÒ%Üœ+×-Ñ-Ó.‰GäÐAÓBÐBà×&Ñ& z°7¸KÓHˆ	à˜IÐ%Ð%r   Nr/   r.   Ú
graph_namec                 óx  — |€|d|› �z   }|g}|€| j                   j                  «       }||v r&||   }|r|j                  dk(  r|j                  d   g}t	        j
                  d||g¬«      }|j                  j                  t	        j                  d|«      g«       | j                   j                  ||¬«       |S )NÚ	_cast_to_ÚCastr   )ÚinputsÚoutputsÚto)r1   )
r   Úoutput_name_to_nodeÚop_typeÚinputr   Ú	make_nodeÚ	attributeÚextendÚmake_attributeÚadd_node)	r   r   r/   r.   r8   r1   r5   Úparent_noder!   s	            r   r,   zFusionUtils.add_cast_node-   sÊ   € ð ÐØ$¨°7°)Ð'<Ñ<ˆKð �ˆØÐ&Ø"&§*¡*×"@Ñ"@Ó"BÐØÐ,Ñ,Ø-¨jÑ9ˆKÙ˜{×2Ñ2°fÒ<Ø%×+Ñ+¨AÑ.Ð/�ä×$Ñ$ V°FÀ[ÀMÔRˆ	à×Ñ×"Ñ"¤F×$9Ñ$9¸$ÀÓ$HÐ#IÔJØ�
‰
×Ñ˜I°*ÐÔ=àÐr   c                 ó&   — | j                  |d«      S )Nr%   )r0   )r   r   s     r   r   zFusionUtils.cast_input_to_int32H   s   € Ø�‰˜z¨7Ó3Ð3r   c                 óœ  — | j                   j                  «       }||   }|D ]¨  }|j                  dk(  sŒd}|j                  D ]<  }|j                  dk(  sŒ|j
                  t        t        j                  «      k(  sŒ:d} n |sŒc|j                  d   }| j                   j                  |«       | j                   j                  ||«       Œª y )Nr4   Fr7   Tr   )r   Úinput_name_to_nodesr9   r<   ÚnameÚir(   r   r   ÚoutputÚremove_nodeÚreplace_input_of_all_nodes)r   r   rC   ÚnodesÚnodeÚis_int32Úattr.   s           r   Úremove_cast_int32zFusionUtils.remove_cast_int32K   s¶   € Ø"Ÿj™j×<Ñ<Ó>ÐØ# JÑ/ˆØò 
	SˆDØ�|‰|˜vÓ%Ø �ØŸ>™>ò �CØ—x‘x 4Ó'¨C¯E©E´S¼×9JÑ9JÓ5KÓ,KØ#'˜Ùðò Ø"&§+¡+¨a¡.�KØ—J‘J×*Ñ*¨4Ô0Ø—J‘J×9Ñ9¸+ÀzÕRñ
	Sr   c                 ó*  — d}| j                   |   |v rP| || j                   |      v r<|| j                   |      j                  | «       t        || j                   |      «      }|| j                   |<   ||v r||   j                  | «       |S | g||<   |S )Nr   )r:   ÚremoveÚlenÚappend)rJ   rE   Únew_input_namerC   Úold_input_references        r   Úupdate_node_inputzFusionUtils.update_node_inputZ   s­   € àÐØ�J‰J�q‰MÐ0Ñ0°dÐ>QÐRV×R\ÑR\Ð]^ÑR_Ñ>`Ñ6`Ø §
¡
¨1¡Ñ.×5Ñ5°dÔ;Ü"%Ð&9¸$¿*¹*ÀQ¹-Ñ&HÓ"IÐà&ˆ�
‰
�1‰àÐ0Ñ0Ø Ñ/×6Ñ6°tÔ<ð #Ð"ð 48°&Ð Ñ/à"Ð"r   c                 ó¤   — |j                   |   }|j                   |   }t        j                  ||||«      }|dk(  xr | j                  |«       }	|	S )a  
        Before:
              (input)-->parent-->node-->(output)
        After:
              (input)-->parent-->
                |
                +----->node-->(output)

        This function returns a flag whether the parent node can be removed.
        r   )r:   r   rT   Úfind_graph_output)
r   rJ   r@   rC   Únode_input_indexÚparent_input_indexÚold_input_namerR   rS   Úparent_can_be_removeds
             r   Úskip_parentzFusionUtils.skip_parentj   se   € ð Ÿ™Ð$4Ñ5ˆØ$×*Ñ*Ð+=Ñ>ˆÜ)×;Ñ;¸DÐBRÐTbÐdwÓxÐð "5¸Ñ!9Ò jÀ5×CZÑCZÐ[iÓCjÐ?jÐà$Ð$r   rJ   c                 ó  — |j                   dv sJ ‚t        |j                  «      dkD  r(| j                  j	                  |j                  d   «      S d }|j
                  D ]'  }|j                  dk(  sŒt        j                  |«      }Œ) |S )N)ÚSqueezeÚ	Unsqueezeé   Úaxes)	r9   rP   r:   r   Úget_constant_valuer<   rD   r   Úget_attribute_value)r   rJ   r`   Úattrs       r   Úget_squeeze_or_unsqueeze_axesz)FusionUtils.get_squeeze_or_unsqueeze_axes€   s   € Ø�|‰|Ð7Ñ7Ð7Ð7ô ˆt�z‰z‹?˜QÒØ—:‘:×0Ñ0°·±¸A±Ó?Ð?àˆØ—N‘Nò 	8ˆDØ�y‰y˜FÓ"Ü×1Ñ1°$Ó7‘ð	8ð ˆr   Úattribute_namec                 óè   — |}| j                   D ]'  }|j                  |k(  sŒt        j                  |«      }Œ) t	        |t
        «      r&t	        |t        t
        f«      xr t        ||d¬«      S ||k(  S )a¦  Verify that a node has expected value for an attribute.

        Args:
            node (NodeProto): a node to check
            attribute_name (str): name of attribute
            expected_value (Any): expected value of the attribute
            default_value (Any, optional): default value if the attribute does not exist. Defaults to None.

        Returns:
            bool: whether the check is passed or not
        F©Ú	equal_nan)r<   rD   r   rb   Ú
isinstanceÚlistr   r   )rJ   re   Úexpected_valueÚdefault_valueÚvaluerc   s         r   Úcheck_node_attributez FusionUtils.check_node_attribute�   sp   € ð ˆØ—N‘Nò 	9ˆDØ�y‰y˜NÓ*Ü×2Ñ2°4Ó8‘ð	9ô �n¤dÔ+Ü˜u¤w´ oÓ6Òo¼KÈÐX]ÐinÔ<oÐoà˜NÑ*Ð*r   Útensorc                 óÚ  — t        | t        «      st        dt        | «      › �«      ‚t	        | j
                  «      dk7  s| j                  t        j                  k7  rt        d«      ‚| j                  rnt        j                  t        j                  | j                  d¬«      | j
                  «      }t        j                  |ddg«      }|j                  «       | _	        | S t        d«      ‚)	z¶Transpose a 2-D INT8 TensorProto
        Args:
            tensor (TensorProto): tensor to be transposed
        Returns:
            tensor (TensorProto): transposed tensor
        z3Expected input type is an ONNX TensorProto but got é   z'Only INT8 2-D tensors can be transposedÚint8)Údtyper_   r   zonly raw buffer supported)ri   r   Ú	TypeErrorr   rP   ÚdimsÚ	data_typeÚINT8r+   Úraw_dataÚnumpyÚreshapeÚ
frombufferÚ	transposeÚtobytes)ro   Ú
int32_dataÚint32_transposed_datas      r   Útranspose_2d_int8_tensorz$FusionUtils.transpose_2d_int8_tensor¤   sÂ   € ô ˜&¤+Ô.ÜÐQÔRVÐW]ÓR^ÐQ_Ð`ÓaÐaäˆv�{‰{Ó˜qÒ  F×$4Ñ$4¼×8HÑ8HÒ$HÜÐFÓGÐGà�?Š?ÜŸ™¤u×'7Ñ'7¸¿¹ÈvÔ'VÐX^×XcÑXcÓdˆJÜ$)§O¡O°JÀÀAÀÓ$GÐ!Ø3×;Ñ;Ó=ˆFŒOð
 ˆô Ð8Ó9Ð9r   c                 óî  — | j                   dvr"t        j                  d| j                   › �«       |j                  | j                  d   «      }|€y|j
                  dk(  xs# |j
                  dk(  xr |j                  d   dk(  }|r|syt        | j                  «      dk(  ry|j                  | j                  d   «      }|j
                  |j
                  k7  ry|€yt        j                  |dk(  «      S )a  Verify if a provided QuantizeLinear (Q) / DequantizeLinear (DQ) node is a good candidate for fusion.
           It is a good candidate for fusion if:
           (1) The Q/DQ node is for per-tensor quantization if allow_per_tensor_quantization_only is `True`
           (2) The Q/DQ node should have constant scale
           (3) The Q/DQ node should have a zero point of 0
        Args:
            node (NodeProto): a Q/DQ node to check
        Returns:
            bool: whether the check is passed or not
        >   ÚQuantizeLinearÚDequantizeLinearz+Provided node is not a Q/DQ node. Op Type: r_   Fr   rq   T)
r9   r   r   ra   r:   ÚndimÚshaperP   ry   Úall)rJ   r   Ú"allow_per_tensor_quantization_onlyÚscaleÚscale_has_single_elementÚ
zero_points         r   Úcheck_qdq_node_for_fusionz%FusionUtils.check_qdq_node_for_fusion¼   sç   € ð �<‰<ÐEÑEÜ�L‰LÐFÀtÇ|Á|ÀnÐUÔVà×(Ñ(¨¯©°A©Ó7ˆð ˆ=Øð $)§:¡:°¡?Ò#_°u·z±zÀQ±Ò7^È5Ï;É;ÐWXÉ>Ð]^ÑK^Ð Ù-Ñ6NØô ˆt�z‰z‹?˜aÒØð ×-Ñ-¨d¯j©j¸©mÓ<ˆ
ð �:‰:˜Ÿ™Ò(Øð ÐØä�y‰y˜ q™Ó)Ð)r   Úinput_indexc                 óü   — t        |j                  «      |kD  sJ ‚| j                  j                  |j                  |   «      }t	        |t
        «      r&t	        |t        t
        f«      xr t        ||d¬«      S ||k(  S )a7  Verify that a node has expected input value

        Args:
            node (NodeProto): a node to check
            input_index (int): index of its input to be verified
            expected_value (Any): expected value of the input

        Returns:
            bool: whether the check is passed or not
        Frg   )rP   r:   r   ra   ri   rj   r   r   )r   rJ   rŒ   rk   rm   s        r   Úcheck_node_input_valuez"FusionUtils.check_node_input_valueç   sm   € ô �4—:‘:‹ Ò,Ð,Ð,à—
‘
×-Ñ-¨d¯j©j¸Ñ.EÓFˆä�n¤dÔ+Ü˜u¤w´ oÓ6Òo¼KÈÐX]ÐinÔ<oÐoà˜NÑ*Ð*r   c                 óÌ  — g }| j                   j                  «       }| j                   j                  «       D ]k  }|j                  dk(  sŒ|j                  d   |vsŒ%| j                   j                  |j                  d   |j                  d   «       |j                  |«       Œm |r>| j                   j                  |«       t        j                  dt        |«      › d�«       yy)z>Remove Identity nodes, except those right before graph output.ÚIdentityr   zRemoved z Identity nodesN)r   Úget_graphs_output_namesrI   r9   rF   rH   r:   rQ   Úremove_nodesr   ÚinforP   )r   Únodes_to_removeÚgraph_output_namesrJ   s       r   Úremove_identity_nodesz!FusionUtils.remove_identity_nodesû   s½   € àˆØ!ŸZ™Z×?Ñ?ÓAÐØ—J‘J×$Ñ$Ó&ò 	1ˆDØ�|‰|˜zÓ)Ø—;‘;˜q‘>Ð);Ò;Ø—J‘J×9Ñ9¸$¿+¹+Àa¹.È$Ï*É*ÐUVÉ-ÔXØ#×*Ñ*¨4Õ0ð		1ñ Ø�J‰J×#Ñ# OÔ4Ü�K‰K˜(¤3 Ó#7Ð"8¸ÐHÕIð r   c                 ó8   — | j                   j                  «        y r   )r   Úremove_cascaded_cast_nodes©r   s    r   r˜   z&FusionUtils.remove_cascaded_cast_nodes	  s   € Ø�
‰
×-Ñ-Õ/r   c                 ó8   — | j                   j                  «        y r   )r   Úremove_useless_cast_nodesr™   s    r   r›   z%FusionUtils.remove_useless_cast_nodes  s   € Ø�
‰
×,Ñ,Õ.r   c                 óH  — | j                   j                  d¬«      }|€yg }| j                   j                  «       D ]�  }|j                  dk(  sŒ|j	                  |j
                  d   «      }|j	                  |j                  d   «      }|sŒR|sŒU||k(  sŒ[t        j                  d|j                  › d|› �«       |j                  |«       Œ’ |�rQt        | j                   j                  «       «      }t        | j                   j                  «       «      }|D �]  }t        t        |j                  «      |z  «      r�t        t        |j
                  «      |z  «      smt        | j                   j!                  «       |j
                  d      «      dk(  r7| j                   j#                  |j
                  d   |j                  d   «       n7Œ´| j                   j%                  |j                  d   |j
                  d   «       | j                   j'                  |«       �Œ yy)	ziRemove reshape node that is not needed based on symbolic shape inference: input and output has same shapeT)ÚupdateNÚReshaper   zRemove reshape node z* since its input shape is same as output: r_   )r   Úinfer_runtime_shaperI   r9   Úget_edge_shaper:   rF   r   r“   rD   rQ   ÚsetÚget_graphs_input_namesr‘   ÚboolrP   rC   Úreplace_output_of_all_nodesrH   rG   )r   Úshape_inferr”   rJ   Úinput_shapeÚoutput_shapeÚgraph_input_namesr•   s           r   Úremove_useless_reshape_nodesz(FusionUtils.remove_useless_reshape_nodes  s»  € à—j‘j×4Ñ4¸DÐ4ÓAˆØÐØàˆØ—J‘J×$Ñ$Ó&ò 	1ˆDØ�|‰|˜yÓ(Ø)×8Ñ8¸¿¹ÀA¹ÓG�Ø*×9Ñ9¸$¿+¹+Àa¹.ÓI�Ú¢<°KÀ<Ó4OÜ—K‘KØ.¨t¯y©y¨kÐ9cÐdoÐcpÐqôð $×*Ñ*¨4Õ0ð	1ò Ü # D§J¡J×$EÑ$EÓ$GÓ HÐÜ!$ T§Z¡Z×%GÑ%GÓ%IÓ!JÐØ'ó -�Üœ˜DŸK™KÓ(Ð+=Ñ=Ô>ä ¤ T§Z¡Z£Ð3DÑ!DÔEÜ §
¡
× >Ñ >Ó @ÀÇÁÈAÁÑ OÓPÐTUÒUàŸ
™
×>Ñ>¸t¿z¹zÈ!¹}ÈdÏkÉkÐZ[ÉnÕ]à à—J‘J×9Ñ9¸$¿+¹+Àa¹.È$Ï*É*ÐUVÉ-ÔXØ—
‘
×&Ñ& tÖ,ñ-ð r   )r%   )NNN)r   r   r   )T)Ú__name__Ú
__module__Ú__qualname__r
   r   ÚstrÚtupler£   r"   r0   r(   r,   r   rM   ÚstaticmethodrT   r[   r   r   rd   rn   r   r€   r‹   rŽ   r–   r˜   r›   r©   © r   r   r   r      sN  „ ð&˜ió &ð!°Cð !¸EÀ$ÈÀ)Ñ<Ló !ñ& Só &ð( #'Ø Ø!%ñàðð ðð ˜4‘Zð	ð ˜$‘Jóð64¨có 4ðS¨Có Sð ñ#ó ð#ð ñ%˜9ò %ó ð%ð*°)ð ÀÈ$Áó ð ñ+°3ò +ó ð+ð, ð¨ò ó ðð. ñ(*¨	ð (*¸)ò (*ó ð(*ðT+¸ó +ò(Jò0ò/ó-r   r   c                   ó,   — e Zd Zeddededefd„«       Zy)ÚNumpyHelperro   Ú
fill_zerosr   c                 ó  — |r4t        | j                  t        j                  | j                  «      ¬«      S | j                  t
        j                  k(  r#dd l}|j                  | «      j                  «       S t        j                  | «      S )N)r…   rs   r   )r   ru   r   Útensor_dtype_to_np_dtyperv   r   ÚBFLOAT16Úonnx_irÚ
from_protory   r	   Úto_array)ro   r³   Úirs      r   r¹   zNumpyHelper.to_array2  ss   € ñ ÜØ—k‘kÜ×5Ñ5°f×6FÑ6FÓGôð ð
 ×Ñœ{×3Ñ3Ò3Û ð —=‘= Ó(×.Ñ.Ó0Ð0Ü×$Ñ$ VÓ,Ð,r   N)F)rª   r«   r¬   r¯   r   r£   r   r¹   r°   r   r   r²   r²   1  s)   „ Øñ-˜ð -°$ð -À7ò -ó ñ-r   r²   )Úloggingr   ry   r   r   Úonnxr   r   r   r	   Ú
onnx_modelr
   rª   r   r   r²   r°   r   r   ú<module>r¾      s:   ðõ
 ã ß &ß =Ó =Ý  á	�8Ó	€÷_-ñ _-÷D	-ò -r   