§
    kŠtj.  ã                  óL   — d dl mZ d dlmZ d dlZddlmZ  G d„ d¦  «        ZdS )é    )Úannotations)ÚdequeNé   )Ú	ONNXModelc                  óú   — e Zd ZdZdCd„ZdDd„ZdEd„Zd„ ZedFd„¦   «         Z	edGd„¦   «         Z
edHd„¦   «         ZedId„¦   «         ZdJd „ZdKdLd%„ZdKdMd&„ZdNd)„Zd*g fdOd/„Zd*d*g d*fdPd5„Z	 	 	 dQdRd9„ZdSd=„Z	 	 dTdUdB„Zd*S )VÚFusionz!
    Base class for fusions.
    Úmodelr   Úfused_op_typeÚstrÚsearch_op_typec                óŽ   — || _         || _        || _        g | _        g | _        | j        dz   | j         z   dz   | _        d | _        d S )NÚ_fused_Ú_)r   r
   r	   Únodes_to_removeÚnodes_to_addÚ_new_node_name_prefixÚ_new_node_name_suffix)Úselfr	   r
   r   s       úe/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/onnxruntime/quantization/fusions/fusion.pyÚ__init__zFusion.__init__   sU   € Ø#1ˆÔØ"/ˆÔØ %ˆŒ
Ø%'ˆÔØ"$ˆÔà%)Ô%7¸)Ñ%CÀdÔFYÑ%YÐ\_Ñ%_ˆÔ"Ø%)ˆÔ"Ð"Ð"ó    Únodeúonnx.NodeProtoÚinput_name_to_nodesúdict[str, list[onnx.NodeProto]]Úoutput_name_to_nodeúdict[str, onnx.NodeProto]c                ó   — t           ‚)z…
        Interface function for derived fusion classes. Tries to fuse a node sequence containing
        the specified node.
        )ÚNotImplementedError)r   r   r   r   s       r   ÚfusezFusion.fuse   s
   € õ "Ð!r   ÚreturnÚboolc                óØ  — | j                              ¦   «         }| j                              ¦   «         }| j                              ¦   «         D ])}|j        | j        k    r|                      |||¦  «         Œ*| j                              | j        ¦  «         | j          	                    | j
        ¦  «         t          | j        p| j
        ¦  «        }|r| j                              ¦   «          |S )z?
        Apply graph fusion on the entire model graph.
        )r	   r   r   ÚnodesÚop_typer   r    Úremove_nodesr   Ú	add_nodesr   r"   Úremove_unused_constant)r   r   r   r   Úgraph_updateds        r   ÚapplyzFusion.apply*   sà   € ð #œj×<Ò<Ñ>Ô>ÐØ"œj×<Ò<Ñ>Ô>Ðà”J×$Ò$Ñ&Ô&ð 	Jð 	JˆDØŒ|˜tÔ2Ò2Ð2Ø—	’	˜$Ð 3Ð5HÑIÔIÐIøàŒ
×Ò Ô 4Ñ5Ô5Ð5ØŒ
×Ò˜TÔ.Ñ/Ô/Ð/å˜TÔ1ÐF°TÔ5FÑGÔGˆàð 	0ØŒJ×-Ò-Ñ/Ô/Ð/àÐr   c                ó    — | 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?   sa   € ØÔ+ˆàÔ%Ð-Ø"&¤*×"IÒ"IÈ&Ñ"QÔ"QˆNØ)7¸!Ñ);ˆDÔ&àÐ<˜dÔ8Ð<Ð<ˆØÐ"Ô" aÑ'Ð"Ô"àˆr   r   úlist[onnx.NodeProto]Úkeep_outputsú	list[str]c                ó^   — | D ])}|j         D ]}||v rŒ||v r||         D ]}|| vr   dS ŒŒ Œ*dS ©NFT)Úoutput)r   r3   r   r   Únode_to_removeÚoutput_to_removeÚimpacted_nodes          r   Úis_safe_to_fuse_nodeszFusion.is_safe_to_fuse_nodesK   s~   € ð .ð 		)ð 		)ˆNØ$2Ô$9ð )ð )Ð Ø# |Ð3Ð3Øà#Ð':Ð:Ð:Ø)<Ð=MÔ)Nð )ð )˜Ø(°Ð?Ð?à#( 5 5 5 5ð @øð)ð ˆtr   Úattribute_namec                óv   — | j         D ]0}|j        |k    r#t          j                             |¦  «        }|c S Œ1d S )N)Ú	attributeÚnameÚonnxÚhelperÚget_attribute_value)r   r<   ÚattrÚvalues       r   Úget_node_attributezFusion.get_node_attribute^   sJ   € à”Nð 	ð 	ˆDØŒy˜NÒ*Ð*Ýœ×7Ò7¸Ñ=Ô=�Ø���ð +ð ˆtr   Únode_outputÚ
child_nodeÚintc                óN   — t          |j        ¦  «        D ]\  }}|| k    r|c S ŒdS )Néÿÿÿÿ)Ú	enumerateÚinput)rF   rG   ÚindexÚ
input_names       r   Úinput_indexzFusion.input_indexf   s?   € å!*¨:Ô+;Ñ!<Ô!<ð 	ð 	ÑˆE�:Ø˜[Ò(Ð(Ø���ð )àˆrr   ú	list[int]c                ó  — g }| j         j        D ]w}|                     d¦  «        r|                     |j        ¦  «         Œ2|                     d¦  «        r|                     |j        ¦  «         Œb|                     d¦  «         Œx|S )NÚ	dim_valueÚ	dim_paramú?)ÚshapeÚdimÚHasFieldÚappendrR   rS   )Útensor_typeÚ
shape_listÚds      r   Útensor_shape_to_listzFusion.tensor_shape_to_listm   s“   € àˆ
ØÔ"Ô&ð 	'ð 	'ˆAØ�zŠz˜+Ñ&Ô&ð 'Ø×!Ò! !¤+Ñ.Ô.Ð.Ð.Ø—’˜KÑ(Ô(ð 'Ø×!Ò! !¤+Ñ.Ô.Ð.Ð.à×!Ò! #Ñ&Ô&Ð&Ð&ØÐr   c                ó~   — t          |j        ¦  «        D ]'\  }}| j                             |¦  «        }|�||fc S Œ(dS )N©NN)rK   rL   r	   Úget_constant_value)r   r   ÚiÚinprD   s        r   Úget_constant_inputzFusion.get_constant_inputy   sS   € Ý ¤
Ñ+Ô+ð 	 ð 	 ‰FˆAˆsØ”J×1Ò1°#Ñ6Ô6ˆEØÐ Ø˜%�x���ð !ð ˆzr   ç�íµ ÷Æ°>Úexpected_valueÚfloatÚdeltac                ó€   — |                       |¦  «        \  }}|�#|j        dk    rt          ||z
  ¦  «        |k     r|S dS )Nr,   rJ   )rb   ÚsizeÚabs)r   r   rd   rf   r`   rD   s         r   Úfind_constant_inputzFusion.find_constant_input�   sK   € Ø×*Ò*¨4Ñ0Ô0‰ˆˆ5ØÐ ¤¨q¢ µS¸ÀÑ9OÑ5PÔ5PÐSXÒ5XÐ5XØˆHàˆrr   c                ó8   — |                       |||¦  «        dk    S ©Nr   )rj   )r   r   rd   rf   s       r   Úhas_constant_inputzFusion.has_constant_inputˆ   s   € Ø×'Ò'¨¨n¸eÑDÔDÈÒIÐIr   Úoutput_nameÚrankc                óv   — | j                              |¦  «        }|€dS t          |j        ¦  «        |k    rdS dS r6   )r	   r_   ÚlenrU   )r   rn   ro   rD   s       r   Úis_constant_with_specified_rankz&Fusion.is_constant_with_specified_rank‹   s@   € Ø”
×-Ò-¨kÑ:Ô:ˆØˆ=Ø�5åˆuŒ{ÑÔ˜tÒ#Ð#Ø�5àˆtr   NÚparent_op_typeú dict[str, onnx.NodeProto] | NoneÚexcludeú(tuple[onnx.NodeProto | None, int | None]c                ó²   — |€| j                              ¦   «         }t          |j        ¦  «        D ]&\  }}||v r||         }|j        |k    r
||vr||fc S Œ'dS )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.
        Nr^   )r	   r   rK   rL   r%   )r   r   rs   r   ru   r`   ra   Úparents           r   Úmatch_first_parentzFusion.match_first_parent•   s~   € ð( Ð&Ø"&¤*×"@Ò"@Ñ"BÔ"BÐå ¤
Ñ+Ô+ð 	%ð 	%‰FˆAˆsØÐ)Ð)Ð)Ø,¨SÔ1�Ø”> ^Ò3Ð3¸ÀgÐ8MÐ8MØ! 1˜9Ð$Ð$Ð$øàˆzr   rO   ú
int | NoneÚreturn_indiceúlist[int] | Noneúonnx.NodeProto | Nonec                óV  — |€J ‚|�|dk    sJ ‚|€| j                              ¦   «         }|€4|                      ||||¦  «        \  }}|�|                     |¦  «         |S |t	          |j        ¦  «        k    rdS | j                              |||¦  «        }|�|j        |k    r||vr|S dS )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   ry   rX   rq   rL   Ú
get_parentr%   )	r   r   rs   rO   r   ru   r{   rx   rM   s	            r   Úmatch_parentzFusion.match_parent´   sÝ   € ð2 ÐÐÐØÐ" k°QÒ&6Ð&6Ð&6Ð6àÐ&Ø"&¤*×"@Ò"@Ñ"BÔ"BÐàÐØ ×3Ò3°D¸.ÐJ]Ð_fÑgÔg‰MˆF�EØÐ(Ø×$Ò$ UÑ+Ô+Ð+ØˆMà�#˜dœj™/œ/Ò)Ð)à�4à”×&Ò& t¨[Ð:MÑNÔNˆØÐ &¤.°NÒ"BÐ"BÀvÐU\ÐG\ÐG\ØˆMàˆtr   Úparent_op_typesÚparent_input_indexúlist[onnx.NodeProto] | Nonec           	     ó8  — |�"t          |¦  «        t          |¦  «        k    sJ ‚|€| j                             ¦   «         }|}g }t          |¦  «        D ]F\  }}	|                      ||	|�||         nd|g |¬¦  «        }
|
€ dS |                     |
¦  «         |
}ŒG|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)ru   r{   )rq   r	   r   rK   r€   rX   )r   r   r�   r‚   r   r{   Úcurrent_nodeÚmatched_parentsr`   r%   Úmatched_parents              r   Úmatch_parent_pathzFusion.match_parent_pathã   s×   € ð0 Ð)ÝÐ)Ñ*Ô*­c°/Ñ.BÔ.BÒBÐBÐBÐBàÐ&Ø"&¤*×"@Ò"@Ñ"BÔ"BÐàˆØˆÝ# OÑ4Ô4ð 	*ð 	*‰JˆAˆwØ!×.Ò.ØØØ);Ð)GÐ" 1Ô%Ð%ÈTØ#ØØ+ð /ñ ô ˆNð Ð%Ø�t�tà×"Ò" >Ñ2Ô2Ð2Ø)ˆLˆLàÐr   Úpathsú!list[tuple[list[str], list[int]]]ú9tuple[int, list[onnx.NodeProto] | None, list[int] | None]c                ó�   — t          |¦  «        D ]5\  }}g }|                      ||d         |d         ||¦  «        }|r|||fc S Œ6dS )z@
        Find a matching parent path to the given node.
        r   r,   )rJ   NN)rK   rˆ   )r   r   r‰   r   r`   Úpathr{   Úmatcheds           r   Úmatch_parent_pathszFusion.match_parent_paths  sn   € õ ! Ñ'Ô'ð 	1ð 	1‰GˆAˆtØˆMØ×,Ò,¨T°4¸´7¸DÀ¼GÐEXÐZgÑhÔhˆGØð 1Ø˜' =Ð0Ð0Ð0Ð0ð1àˆ~r   TÚ
child_typeú&dict[str, list[onnx.NodeProto]] | NoneÚ	recursivec                óV  — | j                              ||¦  «        }t          |¦  «        }t          |¦  «        dk    rk|                     ¦   «         }|j        |k    r|S |r5| j                              ||¦  «        }|D ]}|                     |¦  «         Œt          |¦  «        dk    °kd S rl   )r	   Úget_childrenr   rq   Úpopr%   Ú
appendleft)	r   r   r�   r   r’   ÚchildrenÚdqr…   Úchilds	            r   Úfind_first_child_by_typezFusion.find_first_child_by_type$  s³   € ð ”:×*Ò*¨4Ð1DÑEÔEˆÝ�8‰_Œ_ˆÝ�"‰gŒg˜ŠkˆkØŸ6š6™8œ8ˆLØÔ# zÒ1Ð1Ø#Ð#àð )Øœ:×2Ò2°<ÐATÑUÔU�Ø%ð )ð )�EØ—M’M %Ñ(Ô(Ð(Ð(õ �"‰gŒg˜Škˆkð ˆtr   )r	   r   r
   r   r   r   )r   r   r   r   r   r   )r!   r"   )
r   r2   r3   r4   r   r   r   r   r!   r"   )r   r   r<   r   )rF   r   rG   r   r!   rH   )r!   rP   )r   r   )rc   )r   r   rd   re   rf   re   r!   rH   )r   r   rd   re   rf   re   r!   r"   )rn   r   ro   rH   r!   r"   )
r   r   rs   r   r   rt   ru   r2   r!   rv   )r   r   rs   r   rO   rz   r   rt   ru   r2   r{   r|   r!   r}   )NNN)r   r   r�   r4   r‚   r|   r   rt   r{   r|   r!   rƒ   )r   r   r‰   rŠ   r   r   r!   r‹   )NT)
r   r   r�   r   r   r‘   r’   r"   r!   r}   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r    r*   r1   Ústaticmethodr;   rE   rO   r\   rb   rj   rm   rr   ry   r€   rˆ   r�   rš   © r   r   r   r      sð  € € € € € ðð ð*ð *ð *ð *ð
"ð 
"ð 
"ð 
"ðð ð ð ð*
ð 
ð 
ð ðð ð ñ „\ðð$ ðð ð ñ „\ðð ðð ð ñ „\ðð ð	ð 	ð 	ñ „\ð	ðð ð ð ðð ð ð ð ðJð Jð Jð Jð Jðð ð ð ð AEØ(*ðð ð ð ð ðF #'Ø@DØ(*Ø*.ð-ð -ð -ð -ð -ðf 04Ø@DØ*.ð/ð /ð /ð /ð /ðbð ð ð ð( GKØðð ð ð ð ð ð r   r   )Ú
__future__r   Úcollectionsr   r@   Ú
onnx_modelr   r   r    r   r   ú<module>r¤      s‚   ðð #Ð "Ð "Ð "Ð "Ð "à Ð Ð Ð Ð Ð à €€€à "Ð "Ð "Ð "Ð "Ð "ðhð hð hð hð hñ hô hð hð hð hr   