§
    kŠtjË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dS )
é    )Ú	getLoggerN)Úarray_equalÚndarray)Ú	NodeProtoÚTensorProtoÚhelperÚnumpy_helper)Ú	OnnxModelc            
       óV  — 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dS )%ÚFusionUtilsÚmodelc                 ó   — || _         d S ©N)r   )Úselfr   s     úc/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/onnxruntime/transformers/fusion_utils.pyÚ__init__zFusionUtils.__init__   s   € Ø %ˆŒ
ˆ
ˆ
ó    Ú
input_nameÚreturnc                 ó8  — | j                              |¦  «        }|�Y|j        j        j        t
          j        k    r:|                      |¦  «        \  }}t           	                    d|› d�¦  «         d|fS t           	                    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×1Ò1°*Ñ=Ô=ˆØÐ" {Ô'7Ô'CÔ'MÕQ\ÔQbÒ'bÐ'bØ%)×%=Ò%=¸jÑ%IÔ%IÑ"ˆK˜Ý�LŠLÐD¨zÐDÐDÐDÑEÔEÐEØ˜Ð$Ð$å�ŠÐg°ÐgÐgÈkÐaeÐNeÐgÐgÑhÔhÐhØ�jÐ Ð r   Úint32c                 ó&  — |dz   |z   }|dk    rt          t          j        ¦  «        }nO|dk    rt          t          j        ¦  «        }n/|dk    rt          t          j        ¦  «        }nt          d¦  «        ‚|                      |||¦  «        }||fS )NÚ_r#   Ú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ˆGØ˜IÒ%Ð%Ý�+Ô+Ñ,Ô,ˆGˆGØ˜IÒ%Ð%Ý�+Ô-Ñ.Ô.ˆGˆGåÐAÑBÔBÐBà×&Ò& z°7¸KÑHÔHˆ	à˜IÐ%Ð%r   Nr/   r.   Ú
graph_namec                 óh  — |€|d|› �z   }|g}|€| j                              ¦   «         }||v r#||         }|r|j        dk    r|j        d         g}t	          j        d||g¬¦  «        }|j                             t	          j        d|¦  «        g¦  «         | 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Ô"BÐØÐ,Ð,Ð,Ø-¨jÔ9ˆKØð 0˜{Ô2°fÒ<Ð<Ø%Ô+¨AÔ.Ð/�åÔ$ V°FÀ[ÀMÐRÑRÔRˆ	àÔ×"Ò"¥FÔ$9¸$ÀÑ$HÔ$HÐ#IÑJÔJÐJØŒ
×Ò˜I°*ÐÑ=Ô=Ð=àÐr   c                 ó.   — |                       |d¦  «        S )Nr#   )r0   )r   r   s     r   r   zFusionUtils.cast_input_to_int32H   s   € Ø�Š˜z¨7Ñ3Ô3Ð3r   c                 ój  — | j                              ¦   «         }||         }|D ]Ž}|j        dk    r�d}|j        D ]3}|j        dk    r&|j        t          t          j        ¦  «        k    rd} nŒ4|rB|j	        d         }| j          
                    |¦  «         | j                              ||¦  «         Œ�d S )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Ô/ˆØð 
	Sð 
	SˆDØŒ|˜vÒ%Ð%Ø �Øœ>ð ð �CØ”x 4Ò'Ð'¨C¬EµS½Ô9JÑ5KÔ5KÒ,KÐ,KØ#'˜Ø˜øØð SØ"&¤+¨a¤.�KØ”J×*Ò*¨4Ñ0Ô0Ð0Ø”J×9Ò9¸+ÀzÑRÔRÐRøð
	Sð 
	Sr   c                 ó>  — d}| j         |         |v r[| || j         |                  v rF|| j         |                                       | ¦  «         t          || j         |                  ¦  «        }|| j         |<   ||v r||                              | ¦  «         n| 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�qŒMÐ0Ð0Ð0°dÐ>QÐRVÔR\Ð]^ÔR_Ô>`Ð6`Ð6`Ø ¤
¨1¤Ô.×5Ò5°dÑ;Ô;Ð;Ý"%Ð&9¸$¼*ÀQ¼-Ô&HÑ"IÔ"IÐà&ˆŒ
�1‰àÐ0Ð0Ð0Ø Ô/×6Ò6°tÑ<Ô<Ð<Ð<à37°&Ð Ñ/à"Ð"r   r   c                 ó¬   — |j         |         }|j         |         }t                               ||||¦  «        }|dk    o|                      |¦  «         }	|	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   sf   € ð œÐ$4Ô5ˆØ$Ô*Ð+=Ô>ˆÝ)×;Ò;¸DÐBRÐTbÐdwÑxÔxÐð "5¸Ò!9Ð jÀ5×CZÒCZÐ[iÑCjÔCjÐ?jÐà$Ð$r   rJ   c                 óì   — |j         dv sJ ‚t          |j        ¦  «        dk    r%| j                             |j        d         ¦  «        S d }|j        D ]!}|j        dk    r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Ð7õ ˆtŒz‰?Œ?˜QÒÐØ”:×0Ò0°´¸A´Ñ?Ô?Ð?àˆØ”Nð 	8ð 	8ˆDØŒy˜FÒ"Ð"ÝÔ1°$Ñ7Ô7�øØˆr   Úattribute_namec                 óê   — |}| j         D ]!}|j        |k    rt          j        |¦  «        }Œ"t	          |t
          ¦  «        r.t	          |t          t
          f¦  «        o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�   s   € ð ˆØ”Nð 	9ð 	9ˆDØŒy˜NÒ*Ð*ÝÔ2°4Ñ8Ô8�øå�n¥dÑ+Ô+ð 	+Ý˜u¥wµ oÑ6Ô6Ðo½KÈÐX]ÐinÐ<oÑ<oÔ<oÐoà˜NÒ*Ð*r   Útensorc                 óÚ  — t          | t          ¦  «        st          dt          | ¦  «        › �¦  «        ‚t	          | j        ¦  «        dk    s| j        t          j        k    rt          d¦  «        ‚| j	        rdt          j        t          j        | j	        d¬¦  «        | j        ¦  «        }t          j        |ddg¦  «        }|                     ¦   «         | _	        nt          d¦  «        ‚| S )	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Þ   € õ ˜&¥+Ñ.Ô.ð 	bÝÐ`ÕRVÐW]ÑR^ÔR^Ð`Ð`ÑaÔaÐaåˆvŒ{ÑÔ˜qÒ Ð  FÔ$4½Ô8HÒ$HÐ$HÝÐFÑGÔGÐGàŒ?ð 	:Ýœ¥uÔ'7¸¼ÈvÐ'VÑ'VÔ'VÐX^ÔXcÑdÔdˆJÝ$)¤O°JÀÀAÀÑ$GÔ$GÐ!Ø3×;Ò;Ñ=Ô=ˆFŒOˆOõ Ð8Ñ9Ô9Ð9àˆr   Tc                 óÊ  — | j         dvr"t                               d| j         › �¦  «         |                     | j        d         ¦  «        }|€dS |j        dk    p|j        dk    o|j        d         dk    }|r|sdS t          | j        ¦  «        dk    rdS |                     | j        d         ¦  «        }|j        |j        k    rdS |€dS 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_   NFr   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ÐEÝ�LŠLÐUÀtÄ|ÐUÐUÑVÔVÐVà×(Ò(¨¬°A¬Ñ7Ô7ˆð ˆ=Ø�5ð $)¤:°¢?Ð#_°u´zÀQ²Ð7^È5Ì;ÐWXÌ>Ð]^ÒK^Ð Ø-ð 	Ð6Nð 	Ø�5õ ˆtŒz‰?Œ?˜aÒÐØ�4ð ×-Ò-¨d¬j¸¬mÑ<Ô<ˆ
ð Œ:˜œÒ(Ð(Ø�5ð ÐØ�5åŒy˜ qšÑ)Ô)Ð)r   Úinput_indexc                 ó  — t          |j        ¦  «        |k    sJ ‚| j                             |j        |         ¦  «        }t	          |t
          ¦  «        r.t	          |t          t
          f¦  «        o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ç   s€   € õ �4”:‰Œ Ò,Ð,Ð,Ð,à”
×-Ò-¨d¬j¸Ô.EÑFÔFˆå�n¥dÑ+Ô+ð 	+Ý˜u¥wµ oÑ6Ô6Ðo½KÈÐX]ÐinÐ<oÑ<oÔ<oÐoà˜NÒ*Ð*r   c                 óÆ  — g }| j                              ¦   «         }| j                              ¦   «         D ]b}|j        dk    rU|j        d         |vrF| j                              |j        d         |j        d         ¦  «         |                     |¦  «         Œc|rG| j                              |¦  «         t           
                    dt          |¦  «        › d�¦  «         dS dS )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×?Ò?ÑAÔAÐØ”J×$Ò$Ñ&Ô&ð 	1ð 	1ˆDØŒ|˜zÒ)Ð)Ø”;˜q”>Ð);Ð;Ð;Ø”J×9Ò9¸$¼+Àa¼.È$Ì*ÐUVÌ-ÑXÔXÐXØ#×*Ò*¨4Ñ0Ô0Ð0øàð 	JØŒJ×#Ò# OÑ4Ô4Ð4Ý�KŠKÐH¥3 Ñ#7Ô#7ÐHÐHÐHÑIÔIÐIÐIÐIð	Jð 	Jr   c                 ó8   — | j                              ¦   «          d S r   )r   Úremove_cascaded_cast_nodes©r   s    r   r˜   z&FusionUtils.remove_cascaded_cast_nodes	  s   € ØŒ
×-Ò-Ñ/Ô/Ð/Ð/Ð/r   c                 ó8   — | j                              ¦   «          d S r   )r   Úremove_useless_cast_nodesr™   s    r   r›   z%FusionUtils.remove_useless_cast_nodes  s   € ØŒ
×,Ò,Ñ.Ô.Ð.Ð.Ð.r   c                 óP  — | j                              d¬¦  «        }|€dS g }| j                              ¦   «         D ]‘}|j        dk    r„|                     |j        d         ¦  «        }|                     |j        d         ¦  «        }|rB|r@||k    r:t                               d|j	        › d|› �¦  «         | 
                    |¦  «         Œ’|�rTt          | j                              ¦   «         ¦  «        }t          | j                              ¦   «         ¦  «        }|D �]}t          t          |j        ¦  «        |z  ¦  «        r’t          t          |j        ¦  «        |z  ¦  «        smt          | j                              ¦   «         |j        d                  ¦  «        dk    r2| j                              |j        d         |j        d         ¦  «         n2Œ¹| j                              |j        d         |j        d         ¦  «         | j                              |¦  «         �ŒdS dS )	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×4Ò4¸DÐ4ÑAÔAˆØÐØˆFàˆØ”J×$Ò$Ñ&Ô&ð 	1ð 	1ˆDØŒ|˜yÒ(Ð(Ø)×8Ò8¸¼ÀA¼ÑGÔG�Ø*×9Ò9¸$¼+Àa¼.ÑIÔI�Øð 1 <ð 1°KÀ<Ò4OÐ4OÝ—K’KØq¨t¬yÐqÐqÐdoÐqÐqñô ð ð $×*Ò*¨4Ñ0Ô0Ð0øàñ 	-Ý # D¤J×$EÒ$EÑ$GÔ$GÑ HÔ HÐÝ!$ T¤Z×%GÒ%GÑ%IÔ%IÑ!JÔ!JÐØ'ð -ñ -�Ý�˜DœKÑ(Ô(Ð+=Ñ=Ñ>Ô>ð 	Yå ¥ T¤Z¡¤Ð3DÑ!DÑEÔEð!å ¤
× >Ò >Ñ @Ô @ÀÄÈAÄÔ OÑPÔPÐTUÒUÐUàœ
×>Ò>¸t¼zÈ!¼}ÈdÌkÐZ[ÌnÑ]Ô]Ð]Ð]à à”J×9Ò9¸$¼+Àa¼.È$Ì*ÐUVÌ-ÑXÔXÐ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      sU  € € € € € ð&˜ið &ð &ð &ð &ð!°Cð !¸EÀ$ÈÀ)Ô<Lð !ð !ð !ð !ð&ð & Sð &ð &ð &ð &ð( #'Ø Ø!%ðð àðð ðð ˜4‘Zð	ð ˜$‘Jðð ð ð ð64¨cð 4ð 4ð 4ð 4ðS¨Cð Sð Sð Sð Sð ð#ð #ñ „\ð#ð ð%ð %˜9ð %ð %ð %ñ „\ð%ð*°)ð ÀÈ$Áð ð ð ð ð ð+ð +°3ð +ð +ð +ñ „\ð+ð, ð¨ð ð ð ñ „\ðð. ð(*ð (*¨	ð (*¸)ð (*ð (*ð (*ñ „\ð(*ðT+¸ð +ð +ð +ð +ð(Jð Jð Jð0ð 0ð 0ð/ð /ð /ð-ð -ð -ð -ð -r   r   c                   ó4   — e Zd Zeddededefd„¦   «         ZdS )ÚNumpyHelperFro   Ú
fill_zerosr   c                 ó  — |r-t          | j        t          j        | j        ¦  «        ¬¦  «        S | j        t
          j        k    r+dd l}|                     | ¦  «         	                    ¦   «         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  s‰   € ð ð 	ÝØ”kÝÔ5°fÔ6FÑGÔGðñ ô ð ð
 Ô�{Ô3Ò3Ð3Ø Ð Ð Ð ð —=’= Ñ(Ô(×.Ò.Ñ0Ô0Ð0ÝÔ$ VÑ,Ô,Ð,r   N)F)rª   r«   r¬   r¯   r   r£   r   r¹   r°   r   r   r²   r²   1  sL   € € € € € Øð-ð -˜ð -°$ð -À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   