Ë
    ÝÍ:j.  ã                  óD   — d dl mZ d dlmZ d dlZddlmZ  G d„ d«      Zy)é    )Úannotations)ÚdequeNé   )Ú	ONNXModelc                  ód  — e Zd ZdZdd„Z	 	 	 	 	 	 dd„Zdd„Zd„ Ze	 	 	 	 	 	 	 	 	 	 dd„«       Z	edd„«       Z
edd„«       Zedd	„«       Zdd
„Zddd„Zddd„Zdd„Zdg f	 	 	 	 	 	 	 	 	 d d„Zddg df	 	 	 	 	 	 	 	 	 	 	 	 	 d!d„Z	 	 	 d"	 	 	 	 	 	 	 	 	 	 	 d#d„Z	 	 	 	 	 	 	 	 d$d„Z	 	 d%	 	 	 	 	 	 	 	 	 d&d„Zy)'ÚFusionz!
    Base class for fusions.
    c                ó    — || _         || _        || _        g | _        g | _        | j                  dz   | j                   z   dz   | _        d | _        y )NÚ_fused_Ú_)Úsearch_op_typeÚfused_op_typeÚmodelÚnodes_to_removeÚnodes_to_addÚ_new_node_name_prefixÚ_new_node_name_suffix)Úselfr   r   r   s       ú|/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/onnxruntime/quantization/fusions/fusion.pyÚ__init__zFusion.__init__   sU   € Ø#1ˆÔØ"/ˆÔØ %ˆŒ
Ø%'ˆÔØ"$ˆÔà%)×%7Ñ%7¸)Ñ%CÀd×FYÑFYÑ%YÐ\_Ñ%_ˆÔ"Ø%)ˆÕ"ó    c                ó   — t         ‚)z…
        Interface function for derived fusion classes. Tries to fuse a node sequence containing
        the specified node.
        )ÚNotImplementedError)r   ÚnodeÚinput_name_to_nodesÚoutput_name_to_nodes       r   ÚfusezFusion.fuse   s
   € ô "Ð!r   c                ó  — | j                   j                  «       }| j                   j                  «       }| j                   j                  «       D ]/  }|j                  | j
                  k(  sŒ| j                  |||«       Œ1 | j                   j                  | j                  «       | j                   j                  | j                  «       t        | j                  xs | j                  «      }|r| j                   j                  «        |S )z?
        Apply graph fusion on the entire model graph.
        )r   r   r   ÚnodesÚop_typer   r   Úremove_nodesr   Ú	add_nodesr   ÚboolÚremove_unused_constant)r   r   r   r   Úgraph_updateds        r   ÚapplyzFusion.apply*   sÒ   € ð #Ÿj™j×<Ñ<Ó>ÐØ"Ÿj™j×<Ñ<Ó>Ðà—J‘J×$Ñ$Ó&ò 	JˆDØ�|‰|˜t×2Ñ2Ó2Ø—	‘	˜$Ð 3Ð5HÕIð	Jð 	�
‰
×Ñ × 4Ñ 4Ô5Ø�
‰
×Ñ˜T×.Ñ.Ô/ä˜T×1Ñ1ÒF°T×5FÑ5FÓGˆáØ�J‰J×-Ñ-Ô/àÐr   c                óÊ   — | j                   }| j                  €%| j                  j                  |«      }|dz   | _        |› | j                  ›�}| xj                  dz  c_        |S )Né   )r   r   r   Úget_largest_node_name_suffix)r   ÚprefixÚlargest_suffixÚnew_names       r   Úcreate_unique_node_namezFusion.create_unique_node_name?   se   € Ø×+Ñ+ˆà×%Ñ%Ð-Ø"&§*¡*×"IÑ"IÈ&Ó"QˆNØ)7¸!Ñ);ˆDÔ&à�X˜d×8Ñ8Ð;Ð<ˆØ×"Ò" aÑ'Õ"àˆr   c                ól   — | D ]/  }|j                   D ]  }||v rŒ||v sŒ||   D ]
  }|| vsŒ   y Œ  Œ1 y©NFT)Úoutput)r   Úkeep_outputsr   r   Únode_to_removeÚoutput_to_removeÚimpacted_nodes          r   Úis_safe_to_fuse_nodeszFusion.is_safe_to_fuse_nodesK   sb   € ð .ò 		)ˆNØ$2×$9Ñ$9ò )Ð Ø# |Ñ3Øà#Ð':Ò:Ø)<Ð=MÑ)Nò )˜Ø(°Ò?ã#(ñ)ñ)ð		)ð r   c                óŠ   — | j                   D ]4  }|j                  |k(  sŒt        j                  j	                  |«      }|c S  y )N)Ú	attributeÚnameÚonnxÚhelperÚget_attribute_value)r   Úattribute_nameÚattrÚvalues       r   Úget_node_attributezFusion.get_node_attribute^   s?   € à—N‘Nò 	ˆDØ�y‰y˜NÓ*ÜŸ™×7Ñ7¸Ó=�Ø’ð	ð r   c                óP   — t        |j                  «      D ]  \  }}|| k(  sŒ|c S  y)Néÿÿÿÿ)Ú	enumerateÚinput)Únode_outputÚ
child_nodeÚindexÚ
input_names       r   Úinput_indexzFusion.input_indexf   s3   € ä!*¨:×+;Ñ+;Ó!<ò 	ÑˆE�:Ø˜[Ó(Ø’ð	ð r   c                ó  — g }| j                   j                  D ]m  }|j                  d«      r|j                  |j                  «       Œ0|j                  d«      r|j                  |j
                  «       Œ]|j                  d«       Œo |S )NÚ	dim_valueÚ	dim_paramú?)ÚshapeÚdimÚHasFieldÚappendrI   rJ   )Útensor_typeÚ
shape_listÚds      r   Útensor_shape_to_listzFusion.tensor_shape_to_listm   st   € àˆ
Ø×"Ñ"×&Ñ&ò 	'ˆAØ�z‰z˜+Ô&Ø×!Ñ! !§+¡+Õ.Ø—‘˜KÔ(Ø×!Ñ! !§+¡+Õ.à×!Ñ! #Õ&ð	'ð Ðr   c                ó„   — t        |j                  «      D ](  \  }}| j                  j                  |«      }|€Œ$||fc S  y)N©NN)rA   rB   r   Úget_constant_value)r   r   ÚiÚinpr=   s        r   Úget_constant_inputzFusion.get_constant_inputy   sF   € Ü §
¡
Ó+ò 	 ‰FˆAˆsØ—J‘J×1Ñ1°#Ó6ˆEØÑ Ø˜%�x’ð	 ð
 r   c                ót   — | j                  |«      \  }}|�"|j                  dk(  rt        ||z
  «      |k  r|S y)Nr'   r@   )rY   ÚsizeÚabs)r   r   Úexpected_valueÚdeltarW   r=   s         r   Úfind_constant_inputzFusion.find_constant_input�   s@   € Ø×*Ñ*¨4Ó0‰ˆˆ5ØÐ §¡¨q¢´S¸ÀÑ9OÓ5PÐSXÒ5XØˆHàr   c                ó.   — | j                  |||«      dk\  S ©Nr   )r_   )r   r   r]   r^   s       r   Úhas_constant_inputzFusion.has_constant_inputˆ   s   € Ø×'Ñ'¨¨n¸eÓDÈÑIÐIr   c                ór   — | j                   j                  |«      }|€yt        |j                  «      |k7  ryyr.   )r   rV   ÚlenrL   )r   Úoutput_nameÚrankr=   s       r   Úis_constant_with_specified_rankz&Fusion.is_constant_with_specified_rank‹   s5   € Ø—
‘
×-Ñ-¨kÓ:ˆØˆ=Øäˆu�{‰{Ó˜tÒ#Øàr   Nc                ó¾   — |€| j                   j                  «       }t        |j                  «      D ])  \  }}||v sŒ||   }|j                  |k(  sŒ ||vsŒ%||fc S  y)a  
        Find parent node based on constraints on op_type.

        Args:
            node: current node.
            parent_op_type (str): constraint of parent node op_type.
            output_name_to_node (dict): dictionary with output name as key, and node as value.
            exclude (list): list of nodes that are excluded (not allowed to match as parent).

        Returns:
            parent: The matched parent node. None if not found.
            index: The input index of matched parent node. None if not found.
        rU   )r   r   rA   rB   r   )r   r   Úparent_op_typer   ÚexcluderW   rX   Úparents           r   Úmatch_first_parentzFusion.match_first_parent•   sn   € ð( Ð&Ø"&§*¡*×"@Ñ"@Ó"BÐä §
¡
Ó+ò 	%‰FˆAˆsØÐ)Ò)Ø,¨SÑ1�Ø—>‘> ^Ó3¸ÀgÒ8MØ! 1˜9Ò$ð		%ð r   c                óL  — |€J ‚|�|dk\  sJ ‚|€| j                   j                  «       }|€,| j                  ||||«      \  }}|�|j                  |«       |S |t	        |j
                  «      k\  ry| j                   j                  |||«      }|�|j                  |k(  r||vr|S y)a*  
        Find parent node based on constraints on op_type and index.
        When input_index is None, we will find the first parent node based on constraints,
        and return_indice will be appended the corresponding input index.

        Args:
            node (str): current node name.
            parent_op_type (str): constraint of parent node op_type.
            input_index (int or None): only check the parent given input index of current node.
            output_name_to_node (dict): dictionary with output name as key, and node as value.
            exclude (list): list of nodes that are excluded (not allowed to match as parent).
            return_indice (list): a list to append the input index when input_index is None.

        Returns:
            parent: The matched parent node.
        Nr   )r   r   rl   rO   rd   rB   Ú
get_parentr   )	r   r   ri   rG   r   rj   Úreturn_indicerk   rE   s	            r   Úmatch_parentzFusion.match_parent´   sÄ   € ð2 ÐÐÐØÐ" k°QÒ&6Ð6Ð6àÐ&Ø"&§*¡*×"@Ñ"@Ó"BÐàÐØ ×3Ñ3°D¸.ÐJ]Ð_fÓg‰MˆF�EØÐ(Ø×$Ñ$ UÔ+ØˆMàœ#˜dŸj™j›/Ò)àà—‘×&Ñ& t¨[Ð:MÓNˆØÐ &§.¡.°NÒ"BÀvÐU\ÑG\ØˆMàr   c           	     ó  — |�t        |«      t        |«      k(  sJ ‚|€| j                  j                  «       }|}g }t        |«      D ]:  \  }}	| j	                  ||	|�||   nd|g |¬«      }
|
€ y|j                  |
«       |
}Œ< |S )aJ  
        Find a sequence of input edges based on constraints on parent op_type and index.
        When input_index is None, we will find the first parent node based on constraints,
        and return_indice will be appended the corresponding input index.

        Args:
            node (str): current node name.
            parent_op_types (str): constraint of parent node op_type of each input edge.
            parent_input_index (list): constraint of input index of each input edge. None means no constraint.
            output_name_to_node (dict): dictionary with output name as key, and node as value.
            return_indice (list): a list to append the input index
                                  When there is no constraint on input index of an edge.

        Returns:
            parents: a list of matched parent node.
        N)rj   ro   )rd   r   r   rA   rp   rO   )r   r   Úparent_op_typesÚparent_input_indexr   ro   Úcurrent_nodeÚmatched_parentsrW   r   Úmatched_parents              r   Úmatch_parent_pathzFusion.match_parent_pathã   s¸   € ð0 Ð)ÜÐ)Ó*¬c°/Ó.BÒBÐBÐBàÐ&Ø"&§*¡*×"@Ñ"@Ó"BÐàˆØˆÜ# OÓ4ò 	*‰JˆAˆwØ!×.Ñ.ØØØ);Ð)GÐ" 1Ò%ÈTØ#ØØ+ð /ó ˆNð Ð%Ùà×"Ñ" >Ô2Ø)‰Lð	*ð Ðr   c                óv   — t        |«      D ]+  \  }}g }| j                  ||d   |d   ||«      }|sŒ&|||fc S  y)z@
        Find a matching parent path to the given node.
        r   r'   )r@   NN)rA   rw   )r   r   Úpathsr   rW   Úpathro   Úmatcheds           r   Úmatch_parent_pathszFusion.match_parent_paths  sX   € ô ! Ó'ò 	1‰GˆAˆtØˆMØ×,Ñ,¨T°4¸±7¸DÀ¹GÐEXÐZgÓhˆGÚØ˜' =Ð0Ò0ð		1ð
 r   c                ó:  — | j                   j                  ||«      }t        |«      }t        |«      dkD  rf|j	                  «       }|j
                  |k(  r|S |r4| j                   j                  ||«      }|D ]  }|j                  |«       Œ t        |«      dkD  rŒfy ra   )r   Úget_childrenr   rd   Úpopr   Ú
appendleft)	r   r   Ú
child_typer   Ú	recursiveÚchildrenÚdqrt   Úchilds	            r   Úfind_first_child_by_typezFusion.find_first_child_by_type$  s•   € ð —:‘:×*Ñ*¨4Ð1DÓEˆÜ�8‹_ˆÜ�"‹g˜ŠkØŸ6™6›8ˆLØ×#Ñ# zÒ1Ø#Ð#áØŸ:™:×2Ñ2°<ÐATÓU�Ø%ò )�EØ—M‘M %Õ(ð)ô �"‹g˜‹kð r   )r   r   r   Ústrr   r‡   )r   úonnx.NodeProtor   údict[str, list[onnx.NodeProto]]r   údict[str, onnx.NodeProto])Úreturnr"   )
r   úlist[onnx.NodeProto]r0   ú	list[str]r   r‰   r   rŠ   r‹   r"   )r   rˆ   r;   r‡   )rC   r‡   rD   rˆ   r‹   Úint)r‹   z	list[int])r   rˆ   )g�íµ ÷Æ°>)r   rˆ   r]   Úfloatr^   r�   r‹   rŽ   )r   rˆ   r]   r�   r^   r�   r‹   r"   )re   r‡   rf   rŽ   r‹   r"   )
r   rˆ   ri   r‡   r   ú dict[str, onnx.NodeProto] | Nonerj   rŒ   r‹   z(tuple[onnx.NodeProto | None, int | None])r   rˆ   ri   r‡   rG   z
int | Noner   r�   rj   rŒ   ro   úlist[int] | Noner‹   úonnx.NodeProto | None)NNN)r   rˆ   rr   r�   rs   r‘   r   r�   ro   r‘   r‹   zlist[onnx.NodeProto] | None)r   rˆ   ry   z!list[tuple[list[str], list[int]]]r   rŠ   r‹   z9tuple[int, list[onnx.NodeProto] | None, list[int] | None])NT)
r   rˆ   r�   r‡   r   z&dict[str, list[onnx.NodeProto]] | Noner‚   r"   r‹   r’   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r%   r,   Ústaticmethodr4   r>   rG   rS   rY   r_   rb   rg   rl   rp   rw   r|   r†   © r   r   r   r      s  „ ñó*ð
"àð
"ð =ð
"ð 7ó	
"óò*
ð ðØ-ðàðð =ðð 7ð	ð
 
òó ðð$ òó ðð òó ðð ò	ó ð	óôôJóð AEØ(*ðàðð ðð >ð	ð
 &ðð 
2óðF #'Ø@DØ(*Ø*.ð-àð-ð ð-ð  ð	-ð
 >ð-ð &ð-ð (ð-ð 
ó-ðf 04Ø@DØ*.ð/àð/ð #ð/ð -ð	/ð
 >ð/ð (ð/ð 
%ó/ðbàðð 1ðð 7ð	ð
 
Cóð( GKØðàðð ðð Dð	ð
 ðð 
ôr   r   )Ú
__future__r   Úcollectionsr   r8   Ú
onnx_modelr   r   r˜   r   r   ú<module>rœ      s   ðõ #å ã å "÷hò hr   