§
    kŠtjp  ã                  óN   — d dl mZ d dlZddlmZ ddlmZ  G d„ de¦  «        ZdS )	é    )ÚannotationsNé   )Ú	ONNXModelé   )ÚFusionc                  ó(   ‡ — e Zd Zdˆ fd„Zdd
„Zˆ xZS )ÚFusionLayerNormalizationÚmodelr   c                óN   •— t          ¦   «                              |dd¦  «         d S )NÚLayerNormalizationÚ
ReduceMean)ÚsuperÚ__init__)Úselfr
   Ú	__class__s     €úo/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/onnxruntime/quantization/fusions/fusion_layernorm.pyr   z!FusionLayerNormalization.__init__   s&   ø€ Ý‰Œ×Ò˜Ð 4°lÑCÔCÐCÐCÐCó    Úreduce_mean_nodeúonnx.NodeProtoÚinput_name_to_nodesúdict[str, list[onnx.NodeProto]]Úoutput_name_to_nodeúdict[str, onnx.NodeProto]c           	     óŒ  — | j                              ||¦  «        }t          |¦  «        dk    st          |¦  «        dk    rdS |j        d         }|d         j        dk    s|d         j        d         |k    rdS t          |¦  «        dk    r*|d         j        dk    s|d         j        d         |k    rdS d}|D ]}|                      |d|d¬¦  «        }|� nŒ |€dS |                      |g d	¢g d
¢fg d¢g d¢fg d¢g d
¢fg d¢g d¢fg|¦  «        \  }}	}
|dk     rdS |	d         }||vrdS |	d         }|                      |¦  «        \  }}|�|dk    s|dk    rdS |	d         }|j        dk    r|                      |d¦  «        dk    rdS |j        dk    r|j        d         |j        d         k    rdS ||j	        d                  d         }|j        dk    rdS ||j	        d                  d         }|j        dk    rdS |g}| 
                    |¦  «         | 
                    |	dd…         ¦  «         | 
                    |||g¦  «         |                      ||j	        ||¦  «        sdS |j        d|                      |j	        d         |¦  «        z
           }|                      |d¦  «        sdS |j        d|                      |j	        d         |¦  «        z
           }|                      |d¦  «        sdS | j         
                    |¦  «         t          j                             d|                      ¦   «         |j        d         ||g|j	        d         g¬¦  «        }|j         
                    t          j                             dt+          |¦  «        ¦  «        g¦  «         | j                             |¦  «         dS )a4  
        Interface function that tries to fuse a node sequence containing a ReduceMean node into a single
        LayerNormalization node.

              +----------------------+
              |                      |
              |                      v
          [Root] --> ReduceMean -->  Sub  --> Pow --> ReduceMean --> Add --> Sqrt --> Div --> Mul --> Add
                     (axis=2 or -1)  |      (Y=2)   (axis=2 or -1)  (E-6 or E-12 or 0) ^
                                     |                                                 |
                                     +-------------------------------------------------+

         Or, using Mul instead of Pow:

              +----------------------+
              |                      |
              |                      v
          [Root] --> ReduceMean -->  Sub  --> Mul --> ReduceMean --> Add --> Sqrt --> Div --> Mul --> Add
                     (axis=2 or -1)  |     (in0=in1)   (axis=2 or -1)  (E-6 or E-12 or 0) ^
                                     |                                                 |
                                     +-------------------------------------------------+

         It also handles cases of duplicated sub nodes exported from older version of PyTorch:

              +----------------------+
              |                      v
              |           +-------> Sub-----------------------------------------------+
              |           |                                                           |
              |           |                                                           v
          [Root] --> ReduceMean -->  Sub  --> (Pow or Mul) --> ReduceMean --> Add --> Sqrt --> Div  --> Mul --> Add
              |                      ^
              |                      |
              +----------------------+
        r   r   NÚSubr   ÚDivF)Ú	recursive)ÚSqrtÚAddr   ÚPowr   )r   r   r   r   r   )r   r   r   r    ÚCastr   )r   r   r   r   r   r   )r   r   r   ÚMulr   )r   r   r   r"   r!   r   éÿÿÿÿg-Cëâ6?é   r    g       @r"   r   r   )ÚnameÚinputsÚoutputsÚepsilon)r
   Úget_childrenÚlenÚinputÚop_typeÚfind_first_child_by_typeÚmatch_parent_pathsÚget_constant_inputÚfind_constant_inputÚoutputÚextendÚis_safe_to_fuse_nodesÚinput_indexÚis_constant_with_specified_rankÚnodes_to_removeÚonnxÚhelperÚ	make_nodeÚcreate_unique_node_nameÚ	attributeÚmake_attributeÚfloatÚnodes_to_addÚappend)r   r   r   r   ÚchildrenÚ
root_inputÚdiv_nodeÚchildÚpath_idÚparent_nodesÚ_Úsub_nodeÚsecond_add_nodeÚiÚ
add_weightÚpow_or_mul_nodeÚmul_nodeÚlast_add_nodeÚsubgraph_nodesÚweight_inputÚ
bias_inputÚnormalize_nodes                         r   ÚfusezFusionLayerNormalization.fuse   sf  € ðP ”:×*Ò*Ð+;Ð=PÑQÔQˆÝˆx‰=Œ=˜AÒÐ¥ X¡¤°Ò!2Ð!2ØˆFà%Ô+¨AÔ.ˆ
à�AŒ;Ô %Ò'Ð'¨8°A¬;Ô+<¸QÔ+?À:Ò+MÐ+MØˆFåˆx‰=Œ=˜AÒÐØ˜Œ{Ô" eÒ+Ð+¨x¸¬{Ô/@ÀÔ/CÀzÒ/QÐ/QØ�àˆØð 	ð 	ˆEØ×4Ò4°U¸EÐCVÐbgÐ4ÑhÔhˆHØÐ#Ø�ð $àÐØˆFà#'×#:Ò#:Øà<Ð<Ð<¸o¸o¸oÐNØDÐDÐDÐFXÐFXÐFXÐYØ<Ð<Ð<¸o¸o¸oÐNØDÐDÐDÐFXÐFXÐFXÐYð	ð  ñ	$
ô 	$
Ñ ˆ�˜qð �QŠ;ˆ;ØˆFà Ô#ˆØ˜8Ð#Ð#ØˆFà& qœ/ˆØ×/Ò/°Ñ@Ô@‰ˆˆ:ØÐ ¨q¢ °JÀÒ4GÐ4GàˆFà& qœ/ˆØÔ" eÒ+Ð+°×0HÒ0HÈÐZ]Ñ0^Ô0^ÐbcÒ0cÐ0cØˆFØÔ$¨Ò-Ð-°/Ô2GÈÔ2JÈoÔNcÐdeÔNfÒ2fÐ2fØˆFà& x¤°qÔ'9Ô:¸1Ô=ˆØÔ˜uÒ$Ð$ØˆFà+¨H¬O¸AÔ,>Ô?ÀÔBˆØÔ  EÒ)Ð)ØˆFà*Ð+ˆØ×Ò˜hÑ'Ô'Ð'Ø×Ò˜l¨3¨B¨3Ô/Ñ0Ô0Ð0à×Ò˜}¨h¸ÐAÑBÔBÐBØ×)Ò)ØØÔ ØØñ	
ô 
ð 	ð ˆFà”~ a¨$×*:Ò*:¸8¼?È1Ô;MÈxÑ*XÔ*XÑ&XÔYˆØ×3Ò3°LÀ!ÑDÔDð 	ØˆFà"Ô(¨¨T×-=Ò-=¸h¼oÈaÔ>PÐR_Ñ-`Ô-`Ñ)`Ôaˆ
Ø×3Ò3°JÀÑBÔBð 	ØˆFàÔ×#Ò# NÑ3Ô3Ð3åœ×.Ò.Ø Ø×-Ò-Ñ/Ô/Ø$Ô*¨1Ô-¨|¸ZÐHØ"Ô)¨!Ô,Ð-ð	 /ñ 
ô 
ˆð 	Ô ×'Ò'­¬×)CÒ)CÀIÍuÐU_ÑO`ÔO`Ñ)aÔ)aÐ(bÑcÔcÐcØÔ× Ò  Ñ0Ô0Ð0Ð0Ð0r   )r
   r   )r   r   r   r   r   r   )Ú__name__Ú
__module__Ú__qualname__r   rR   Ú__classcell__)r   s   @r   r	   r	      s_   ø€ € € € € ðDð Dð Dð Dð Dð Dð@1ð @1ð @1ð @1ð @1ð @1ð @1ð @1r   r	   )Ú
__future__r   r7   Ú
onnx_modelr   Úfusionr   r	   © r   r   ú<module>r[      s„   ðð #Ð "Ð "Ð "Ð "Ð "à €€€à "Ð "Ð "Ð "Ð "Ð "Ø Ð Ð Ð Ð Ð ðD1ð D1ð D1ð D1ð D1˜vñ D1ô D1ð D1ð D1ð D1r   