§
    kŠtj­  ã                   ó‚   — d dl mZ d dlmZ d dlmZ d dlZd dlZd dlm	Z	 d dl
mZ  ee¦  «        Z G d„ d¦  «        ZdS )	é    )ÚSequence)Ú	getLogger)ÚAnyN)Úhelper)Ú	OnnxModelc                   óÌ   — e Zd ZdZdej        fd„Zdeddfd„Zde	ddfd	„Z
de	d
ededdfd„Zdd„Zdd„Zdde	dedee         dedef
d„Zddeddfd„Zdd„Zedd„¦   «         ZdS )ÚDynamoOnnxHelperzK
    Helper class for processing ONNX models exported by Torch Dynamo.
    Úmodelc                 ó.   — t          |¦  «        | _        d S )N)r   r
   )Úselfr
   s     úi/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/onnxruntime/transformers/dynamo_onnx_helper.pyÚ__init__zDynamoOnnxHelper.__init__   s   € Ý˜uÑ%Ô%ˆŒ
ˆ
ˆ
ó    Úedge_mappingÚreturnNc                 ó@  — | j         j         j        j        D ]ž}t          t	          |j        ¦  «        ¦  «        D ],}|j        |         |v r||j        |                  |j        |<   Œ-t          t	          |j        ¦  «        ¦  «        D ],}|j        |         |v r||j        |                  |j        |<   Œ-ŒŸ| j         j         j        j        D ]}|j        |v r||j                 |_        Œ| j         j         j        j        D ]}|j        |v r||j                 |_        ŒdS )zP
        Updates the edges in the model according to the given mapping.
        N)r
   ÚgraphÚnodeÚrangeÚlenÚinputÚoutputÚname)r   r   r   ÚiÚgraph_inputÚgraph_outputs         r   Úupdate_edgeszDynamoOnnxHelper.update_edges   s>  € ð ”JÔ$Ô*Ô/ð 	Bð 	BˆDÝ�3˜tœz™?œ?Ñ+Ô+ð @ð @�Ø”:˜a”= LÐ0Ð0Ø$0°´¸A´Ô$?�D”J˜q‘MøÝ�3˜tœ{Ñ+Ô+Ñ,Ô,ð Bð B�Ø”;˜q”> \Ð1Ð1Ø%1°$´+¸a´.Ô%A�D”K ‘NøðBð  œ:Ô+Ô1Ô7ð 	Bð 	BˆKØÔ <Ð/Ð/Ø#/°Ô0@Ô#A�Ô øØ œJÔ,Ô2Ô9ð 	Dð 	DˆLØÔ  LÐ0Ð0Ø$0°Ô1BÔ$C�Ô!øð	Dð 	Dr   Ú	func_namec                 ó.  — t                                d|› d�¦  «         g }g }g }g }| j        j        j        j        D ]^}|j        |k    rQ|                     |¦  «         |                     t          |j	        ¦  «        t          |j
        ¦  «        z   ¦  «         Œ_d}| j        j        j        D ]r}|j        |k    re|                     t          |j        ¦  «        ¦  «         |                     t          |j	        ¦  «        t          |j
        ¦  «        z   ¦  «         |}Œst          |¦  «        t          |¦  «        k    sJ ‚|D ]+}| j        j        j        j                             |¦  «         Œ,|D ]+}| j        j        j        j                             |¦  «         Œ,|�$| j        j        j                             |¦  «         i }	t          t          |¦  «        ¦  «        D ]}
||
         }||
         }||k    r||	|<   Œ|                      |	¦  «        S )zH
        Unrolls the function with the given name in the model.
        zUnrolling function z...N)ÚloggerÚdebugr
   r   r   Úop_typeÚappendÚextendÚlistr   r   Ú	functionsr   r   Úremover   r   )r   r   Únodes_to_removeÚnodes_to_addÚedges_to_removeÚedges_to_addr   Úfunc_to_removeÚfr   r   ÚkÚvs                r   Úunroll_functionz DynamoOnnxHelper.unroll_function,   s  € õ 	�ŠÐ9¨9Ð9Ð9Ð9Ñ:Ô:Ð:ØˆØˆØˆØˆØ”JÔ$Ô*Ô/ð 	Mð 	MˆDØŒ|˜yÒ(Ð(Ø×&Ò& tÑ,Ô,Ð,Ø×&Ò&¥t¨D¬JÑ'7Ô'7½$¸t¼{Ñ:KÔ:KÑ'KÑLÔLÐLøàˆØ”Ô!Ô+ð 	#ð 	#ˆAØŒv˜Ò"Ð"Ø×#Ò#¥D¨¬¡L¤LÑ1Ô1Ð1Ø×#Ò#¥D¨¬¡M¤MµD¸¼±N´NÑ$BÑCÔCÐCØ!"�øå�?Ñ#Ô#¥s¨<Ñ'8Ô'8Ò8Ð8Ð8Ð8à#ð 	5ð 	5ˆDØŒJÔÔ"Ô'×.Ò.¨tÑ4Ô4Ð4Ð4Ø ð 	5ð 	5ˆDØŒJÔÔ"Ô'×.Ò.¨tÑ4Ô4Ð4Ð4ØÐ%ØŒJÔÔ&×-Ò-¨nÑ=Ô=Ð=àˆÝ•s˜?Ñ+Ô+Ñ,Ô,ð 	$ð 	$ˆAØ Ô"ˆAØ˜Q”ˆAØ�AŠvˆvØ"#�˜Q‘øà× Ò  Ñ.Ô.Ð.r   Úinput_idÚ	output_idc                 ób  — i }g }| j         j         j        j        D ]P}|j                             |¦  «        dk    r0|j        |         ||j        |         <   |                     |¦  «         ŒQ|D ]+}| j         j         j        j                             |¦  «         Œ,|  	                    |¦  «         dS )z4
        Removes the function in the model.
        éÿÿÿÿN)
r
   r   r   r"   Úfindr   r   r#   r'   r   )r   r   r1   r2   r   r(   r   s          r   Úremove_functionz DynamoOnnxHelper.remove_functionS   s»   € ð ˆØˆØ”JÔ$Ô*Ô/ð 	-ð 	-ˆDØŒ|× Ò  Ñ+Ô+¨rÒ1Ð1Ø59´[ÀÔ5K�˜TœZ¨Ô1Ñ2Ø×&Ò& tÑ,Ô,Ð,øØ#ð 	5ð 	5ˆDØŒJÔÔ"Ô'×.Ò.¨tÑ4Ô4Ð4Ð4à×Ò˜,Ñ'Ô'Ð'Ð'Ð'r   c                 óh   — t                                d¦  «         |                      ddd¦  «         dS )z9
        Removes the dropout layer in the model.
        zRemoving dropout layer...ÚDropoutr   N©r    r!   r6   ©r   s    r   Úremove_dropout_layerz%DynamoOnnxHelper.remove_dropout_layerb   s5   € õ 	�ŠÐ0Ñ1Ô1Ð1Ø×Ò˜Y¨¨1Ñ-Ô-Ð-Ð-Ð-r   c                 óh   — t                                d¦  «         |                      ddd¦  «         dS )z9
        Removes the LM head layer in the model.
        zRemoving LM head layer...ÚLinear_lm_headé   r   Nr9   r:   s    r   Úremove_lm_head_layerz%DynamoOnnxHelper.remove_lm_head_layeri   s6   € õ 	�ŠÐ0Ñ1Ô1Ð1à×ÒÐ-¨q°!Ñ4Ô4Ð4Ð4Ð4r   Tr   Ú	data_typeÚdimsÚvalsÚrawc                 ó   — |r˜t          j        |¦  «        }t          |t          j        ¦  «        s)t          j        ||¬¦  «                             ¦   «         }n'|                     |¦  «                             ¦   «         }t          j        ||||d¬¦  «        }nt          j        ||||d¬¦  «        }| j	         
                    |¦  «         |S )N)ÚdtypeT)r   r@   rA   rB   rC   F)r   Útensor_dtype_to_np_dtypeÚ
isinstanceÚnpÚndarrayÚarrayÚtobytesÚastypeÚmake_tensorr
   Úadd_initializer)	r   r   r@   rA   rB   rC   Únp_typeÚbytesÚtensors	            r   rN   z DynamoOnnxHelper.add_initializerq   sÚ   € Øð 	ÝÔ5°iÑ@Ô@ˆGÝ˜d¥B¤JÑ/Ô/ð 7Ýœ ¨WÐ5Ñ5Ô5×=Ò=Ñ?Ô?��àŸš GÑ,Ô,×4Ò4Ñ6Ô6�ÝÔ'ØØ#ØØØðñ ô ˆFˆFõ Ô'ØØ#ØØØðñ ô ˆFð 	Œ
×"Ò" 6Ñ*Ô*Ð*Øˆr   é   Úmin_sizec           	      óö  — t                                d|› d�¦  «         | j                             d¦  «        }g }|D ]¡}| j                             |j        d         ¦  «        }|�|j        |k     rŒ5|j        D ]O}|j        dk    rB|  	                    |j        d         |j
        j        t          |j        ¦  «        |¬¦  «          nŒP|                     |¦  «         Œ¢| j                             |¦  «         dS )zT
        Converts Constant ops of size [min_size] or higher to initializers
        z'Converting constants greater than size z to initializersÚConstantr   NÚvalue)r   r@   rA   rB   )r    r!   r
   Úget_nodes_by_op_typeÚget_constant_valuer   ÚsizeÚ	attributer   rN   Útr@   r%   Úshaper#   Úremove_nodes)r   rS   Úconstant_nodesr(   r   Únp_dataÚatts          r   Ú!convert_constants_to_initializersz2DynamoOnnxHelper.convert_constants_to_initializers‹   s  € õ 	�ŠÐY¸xÐYÐYÐYÑZÔZÐZàœ×8Ò8¸ÑDÔDˆØˆà"ð 	)ð 	)ˆDà”j×3Ò3°D´KÀ´NÑCÔCˆGð ˆ '¤,°Ò"9Ð"9Øð ”~ð ð �Ø”8˜wÒ&Ð&Ø×(Ò(Ø!œ[¨œ^Ø"%¤%¤/Ý! '¤-Ñ0Ô0Ø$ð	 )ñ ô ð ð �Eð 'ð ×"Ò" 4Ñ(Ô(Ð(Ð(ð 	Œ
×Ò Ñ0Ô0Ð0Ð0Ð0r   c                 óÊ   — | j                              ¦   «         D ]}|                     d¦  «         Œ| j                              ¦   «         D ]}|                     d¦  «         ŒdS )z4
        Clear metadata fields in all nodes
        Úmetadata_propsN)r
   ÚgraphsÚ
ClearFieldÚnodes)r   r   r   s      r   Úclear_metadatazDynamoOnnxHelper.clear_metadata¬   sv   € ð ”Z×&Ò&Ñ(Ô(ð 	/ð 	/ˆEØ×ÒÐ-Ñ.Ô.Ð.Ð.Ø”J×$Ò$Ñ&Ô&ð 	.ð 	.ˆDØ�OŠOÐ,Ñ-Ô-Ð-Ð-ð	.ð 	.r   c                 óR  — ddl m} | j        j                             ¦   «         D �]€\  }}|                     ¦   «         }t          |¦  «        dk    �rR|d         j        dk    �r@|d         }|j         	                    d¦  «        }|€?| 
                    |j                             ¦   «                              ¦   «         ¦  «        }nQ| 
                    |j                             ¦   «                              |                     ¦   «         ¦  «        ¦  «        }|                     |j        |j        |                     |j        ¦  «        |¬¦  «        }|j                             |j        d         |¦  «         || j        j        |<   |j                             |d¬	¦  «         �Œ‚dS )
z]
        Constant fold Transpose initializers without changing the initializer names
        r   )ÚirrR   Ú	TransposeÚpermN)r   r\   ÚtypeÚconst_valueT)Úsafe)Ú
onnxscriptri   r   ÚinitializersÚitemsÚ	consumersr   r"   Ú
attributesÚgetrQ   rm   ÚnumpyÚ	transposeÚas_intsÚValuer   r\   Ú
TensorTyperE   ÚconvenienceÚreplace_all_uses_withÚoutputsr'   )	r
   ri   r   ÚinitializerÚ
user_nodesÚtranspose_noderk   Útransposed_tensorÚnew_initializers	            r   Úfold_transpose_initializersz,DynamoOnnxHelper.fold_transpose_initializersµ   s•  € ð
 	"Ð!Ð!Ð!Ð!Ð!à!&¤Ô!9×!?Ò!?Ñ!AÔ!Að 	Gñ 	GÑˆD�+Ø$×.Ò.Ñ0Ô0ˆJÝ�:‰Œ !Ò#Ñ#¨
°1¬Ô(=ÀÒ(LÑ(LØ!+¨A¤�Ø%Ô0×4Ò4°VÑ<Ô<�Ø�<Ø(*¯	ª	°+Ô2I×2OÒ2OÑ2QÔ2Q×2[Ò2[Ñ2]Ô2]Ñ(^Ô(^Ð%Ð%à(*¯	ª	°+Ô2I×2OÒ2OÑ2QÔ2Q×2[Ò2[Ð\`×\hÒ\hÑ\jÔ\jÑ2kÔ2kÑ(lÔ(lÐ%Ø"$§(¢(Ø$Ô)Ø+Ô1ØŸšÐ'8Ô'>Ñ?Ô?Ø 1ð	 #+ñ #ô #�ð ”×4Ò4°^Ô5KÈAÔ5NÐP_Ñ`Ô`Ð`Ø1@�”Ô(¨Ñ.ØÔ$×+Ò+¨NÀÐ+ÑFÔFÐFùð#	Gð 	Gr   )r   N)T)rR   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚonnxÚ
ModelProtor   Údictr   Ústrr0   Úintr6   r;   r?   r   r   ÚboolrN   ra   rg   Ústaticmethodr‚   © r   r   r	   r	      sƒ  € € € € € ðð ð&˜dœoð &ð &ð &ð &ðD¨ð D°$ð Dð Dð Dð Dð&%/¨ð %/°ð %/ð %/ð %/ð %/ðN(¨ð (¸ð (Èð (ÐPTð (ð (ð (ð (ð.ð .ð .ð .ð5ð 5ð 5ð 5ðð  Cð °Cð ¸xÈ¼}ð ÐTWð Ð^bð ð ð ð ð41ð 1¸#ð 1Àdð 1ð 1ð 1ð 1ðB.ð .ð .ð .ð ðGð Gð Gñ „\ðGð Gð Gr   r	   )Úcollections.abcr   Úloggingr   Útypingr   ru   rH   r‡   r   Ú
onnx_modelr   rƒ   r    r	   rŽ   r   r   ú<module>r“      sË   ðð
 %Ð $Ð $Ð $Ð $Ð $Ø Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð à Ð Ð Ð Ø €€€Ø Ð Ð Ð Ð Ð Ø  Ð  Ð  Ð  Ð  Ð  à	ˆ�8Ñ	Ô	€ð|Gð |Gð |Gð |Gð |Gñ |Gô |Gð |Gð |Gð |Gr   