Ë
    �…j­  ã                   ót   — 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y)	é    )Ú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y)ÚDynamoOnnxHelperzK
    Helper class for processing ONNX models exported by Torch Dynamo.
    Úmodelc                 ó$   — t        |«      | _        y )N)r   r
   )Úselfr
   s     úl/root/aria/tools/markitdown-venv/lib/python3.12/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 ]3  }|j
                  |   |v sŒ||j
                  |      |j
                  |<   Œ5 t        t	        |j                  «      «      D ]3  }|j                  |   |v sŒ||j                  |      |j                  |<   Œ5 Œ¬ | j                   j                   j                  j
                  D ]%  }|j                  |v sŒ||j                     |_        Œ' | j                   j                   j                  j                  D ]%  }|j                  |v sŒ||j                     |_        Œ' y)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   sR  € ð —J‘J×$Ñ$×*Ñ*×/Ñ/ò 	BˆDÜœ3˜tŸz™z›?Ó+ò @�Ø—:‘:˜a‘= LÒ0Ø$0°·±¸A±Ñ$?�D—J‘J˜q’Mð@ô œ3˜tŸ{™{Ó+Ó,ò B�Ø—;‘;˜q‘> \Ò1Ø%1°$·+±+¸a±.Ñ%A�D—K‘K ’NñBð		Bð  Ÿ:™:×+Ñ+×1Ñ1×7Ñ7ò 	BˆKØ×Ñ <Ò/Ø#/°×0@Ñ0@Ñ#A�Õ ð	Bð !ŸJ™J×,Ñ,×2Ñ2×9Ñ9ò 	DˆLØ× Ñ  LÒ0Ø$0°×1BÑ1BÑ$C�Õ!ñ	Dr   Ú	func_namec                 óœ  — t         j                  d|› d�«       g }g }g }g }| j                  j                  j                  j                  D ]]  }|j
                  |k(  sŒ|j                  |«       |j                  t        |j                  «      t        |j                  «      z   «       Œ_ d}| j                  j                  j                  D ]r  }|j                  |k(  sŒ|j                  t        |j                  «      «       |j                  t        |j                  «      t        |j                  «      z   «       |}Œt t        |«      t        |«      k(  sJ ‚|D ];  }| j                  j                  j                  j                  j                  |«       Œ= |D ];  }| j                  j                  j                  j                  j                  |«       Œ= |�/| j                  j                  j                  j                  |«       i }	t        t        |«      «      D ]  }
||
   }||
   }||k7  sŒ||	|<   Œ | j!                  |	«      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¨+°SÐ9Ô:ØˆØˆØˆØˆØ—J‘J×$Ñ$×*Ñ*×/Ñ/ò 	MˆDØ�|‰|˜yÓ(Ø×&Ñ& tÔ,Ø×&Ñ&¤t¨D¯J©JÓ'7¼$¸t¿{¹{Ó:KÑ'KÕLð	Mð
 ˆØ—‘×!Ñ!×+Ñ+ò 	#ˆAØ�v‰v˜Ó"Ø×#Ñ#¤D¨¯©£LÔ1Ø×#Ñ#¤D¨¯©£M´D¸¿¹³NÑ$BÔCØ!"‘ð		#ô �?Ó#¤s¨<Ó'8Ò8Ð8Ð8à#ò 	5ˆDØ�J‰J×Ñ×"Ñ"×'Ñ'×.Ñ.¨tÕ4ð	5à ò 	5ˆDØ�J‰J×Ñ×"Ñ"×'Ñ'×.Ñ.¨tÕ4ð	5àÐ%Ø�J‰J×Ñ×&Ñ&×-Ñ-¨nÔ=àˆÜ”s˜?Ó+Ó,ò 	$ˆAØ Ñ"ˆAØ˜Q‘ˆAØ�A‹vØ"#�˜Q’ð		$ð × Ñ  Ó.Ð.r   Úinput_idÚ	output_idc                 óª  — i }g }| j                   j                   j                  j                  D ]Q  }|j                  j	                  |«      dk7  sŒ"|j
                  |   ||j                  |   <   |j                  |«       ŒS |D ];  }| j                   j                   j                  j                  j                  |«       Œ= | j                  |«       y)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‘J×$Ñ$×*Ñ*×/Ñ/ò 	-ˆDØ�|‰|× Ñ  Ó+¨rÓ1Ø59·[±[ÀÑ5K�˜TŸZ™Z¨Ñ1Ñ2Ø×&Ñ& tÕ,ð	-ð $ò 	5ˆDØ�J‰J×Ñ×"Ñ"×'Ñ'×.Ñ.¨tÕ4ð	5ð 	×Ñ˜,Õ'r   c                 óT   — t         j                  d«       | j                  ddd«       y)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   s#   € ô 	�‰Ð0Ô1Ø×Ñ˜Y¨¨1Õ-r   c                 óT   — t         j                  d«       | j                  ddd«       y)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   s$   € ô 	�‰Ð0Ô1à×ÑÐ-¨q°!Õ4r   r   Ú	data_typeÚdimsÚvalsÚrawc                 ó’  — |r�t        j                  |«      }t        |t        j                  «      s&t        j
                  ||¬«      j                  «       }n|j                  |«      j                  «       }t        j                  ||||d¬«      }nt        j                  ||||d¬«      }| j                  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Ñ5°iÓ@ˆGÜ˜d¤B§J¡JÔ/ÜŸ™ ¨WÔ5×=Ñ=Ó?‘àŸ™ GÓ,×4Ñ4Ó6�Ü×'Ñ'ØØ#ØØØô‰Fô ×'Ñ'ØØ#ØØØôˆFð 	�
‰
×"Ñ" 6Ô*Øˆr   Úmin_sizec           	      ó   — t         j                  d|› d�«       | j                  j                  d«      }g }|D ]¸  }| j                  j	                  |j
                  d   «      }|�|j                  |k  rŒ=|j                  D ]\  }|j                  dk(  sŒ| j                  |j
                  d   |j                  j                  t        |j                  «      |¬«        n |j                  |«       Œº | j                  j                  |«       y)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   rR   Úconstant_nodesr(   r   Únp_dataÚatts          r   Ú!convert_constants_to_initializersz2DynamoOnnxHelper.convert_constants_to_initializers‹   sö   € ô 	�‰Ð>¸x¸jÐHXÐYÔZàŸ™×8Ñ8¸ÓDˆØˆà"ò 	)ˆDà—j‘j×3Ñ3°D·K±KÀ±NÓCˆGð ˆ '§,¡,°Ò"9Øð —~‘~ò �Ø—8‘8˜wÓ&Ø×(Ñ(Ø!Ÿ[™[¨™^Ø"%§%¡%§/¡/Ü! '§-¡-Ó0Ø$ð	 )ô ñ ðð ×"Ñ" 4Õ(ð'	)ð, 	�
‰
×Ñ Õ0r   c                 óÄ   — | j                   j                  «       D ]  }|j                  d«       Œ | j                   j                  «       D ]  }|j                  d«       Œ y)z4
        Clear metadata fields in all nodes
        Úmetadata_propsN)r
   ÚgraphsÚ
ClearFieldÚnodes)r   r   r   s      r   Úclear_metadatazDynamoOnnxHelper.clear_metadata¬   sX   € ð —Z‘Z×&Ñ&Ó(ò 	/ˆEØ×ÑÐ-Õ.ð	/à—J‘J×$Ñ$Ó&ò 	.ˆDØ�O‰OÐ,Õ-ñ	.r   c                 óR  — ddl m} | j                  j                  j	                  «       D �]y  \  }}|j                  «       }t        |«      dk(  sŒ&|d   j                  dk(  sŒ9|d   }|j                  j                  d«      }|€8|j                  |j                  j                  «       j                  «       «      }nF|j                  |j                  j                  «       j                  |j                  «       «      «      }|j                  |j                   |j"                  |j%                  |j&                  «      |¬«      }|j(                  j+                  |j,                  d   |«       || j                  j                  |<   |j                  j/                  |d¬	«       �Œ| y)
z]
        Constant fold Transpose initializers without changing the initializer names
        r   )Úiré   Ú	TransposeÚpermN)r   r[   ÚtypeÚconst_valueT)Úsafe)Ú
onnxscriptrh   r   ÚinitializersÚitemsÚ	consumersr   r"   Ú
attributesÚgetrQ   rm   ÚnumpyÚ	transposeÚas_intsÚValuer   r[   Ú
TensorTyperE   ÚconvenienceÚreplace_all_uses_withÚoutputsr'   )	r
   rh   r   ÚinitializerÚ
user_nodesÚtranspose_noderk   Útransposed_tensorÚnew_initializers	            r   Úfold_transpose_initializersz,DynamoOnnxHelper.fold_transpose_initializersµ   sc  € õ
 	"à!&§¡×!9Ñ!9×!?Ñ!?Ó!Aó 	GÑˆD�+Ø$×.Ñ.Ó0ˆJÜ�:‹ !Ó#¨
°1©×(=Ñ(=ÀÓ(LØ!+¨A¡�Ø%×0Ñ0×4Ñ4°VÓ<�Ø�<Ø(*¯	©	°+×2IÑ2I×2OÑ2OÓ2Q×2[Ñ2[Ó2]Ó(^Ñ%à(*¯	©	°+×2IÑ2I×2OÑ2OÓ2Q×2[Ñ2[Ð\`×\hÑ\hÓ\jÓ2kÓ(lÐ%Ø"$§(¡(Ø$×)Ñ)Ø+×1Ñ1ØŸ™Ð'8×'>Ñ'>Ó?Ø 1ð	 #+ó #�ð —‘×4Ñ4°^×5KÑ5KÈAÑ5NÐP_Ô`Ø1@�—‘×(Ñ(¨Ñ.Ø×$Ñ$×+Ñ+¨NÀÐ+ÖFñ#	Gr   )r   N)T)ri   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚonnxÚ
ModelProtor   Údictr   Ústrr0   Úintr6   r;   r?   r   r   ÚboolrN   r`   rf   Ústaticmethodr‚   © r   r   r	   r	      sÊ   „ ñð&˜dŸo™oó &ðD¨ð D°$ó Dð&%/¨ð %/°ó %/ðN(¨ð (¸ð (Èð (ÐPTó (ó.ó5ñ Cð °Cð ¸xÈ¹}ð ÐTWð Ð^bó ñ41¸#ð 1Àdó 1óB.ð òGó ñGr   r	   )Úcollections.abcr   Úloggingr   Útypingr   ru   rH   r‡   r   Ú
onnx_modelr   rƒ   r    r	   rŽ   r   r   ú<module>r“      s4   ðõ
 %Ý Ý ã Û Ý Ý  á	�8Ó	€÷|Gò |Gr   