Ë
    g^(h)  ã                  óÒ  — d dl mZ d dlZd dlZd dlZd dlmZmZ d dlZd dl	m
c mZ d dlmZmZ d dlmZ d dlmZ d dlmZmZ d dlmZmZ d d	lmZ d d
lmZ ddlmZ ej@                  jB                  Z!ej@                  jD                  Z" ejF                  e$«      Z%	 	 	 	 	 	 	 	 dd„Z&	 	 	 	 	 	 	 	 dd„Z'	 	 	 	 	 	 	 	 dd„Z( G d„ dejR                  «      Z*d„ Z+dd„Z,dd„Z-dd„Z.dd„Z/y)é    )ÚannotationsN)ÚAnyÚOptional)Údynamo_timedÚlazy_format_graph_code)ÚMutationType)Úfx_graph_cse)Úconstant_foldÚreplace_node_with_constant)Úenter_freezingÚrecord_has_frozen_params)Úfreezing_passes)Úview_to_reshapeé   )Úconfigc                ó€  — | j                   j                  d¬«      }|dt        |«       }g }|j                  D �cg c]  }|j                  �|j                  ‘Œ }}t        |j                  «      D ��	cg c]3  \  }}	|	j                  t        j                  t        j                  fv r|‘Œ5 }
}}	t        t        ||«      «      D ]/  \  }\  }}||
v s||v r|j                  |«       Œ#t        | ||«       Œ1 |j                  t        t        |«      t        |«      «      «       | j!                  «        |S c c}w c c}	}w )zÂ
    Replaces the parameters of a PyTorch GraphModule with constants wherever possible.
    Returns a list of indices representing the input parameters that were not converted to constants.
    Úplaceholder©ÚopN)ÚgraphÚ
find_nodesÚlenÚoutput_infoÚbase_idxÚ	enumerateÚ
input_infoÚmutation_typer   ÚMUTATED_IN_GRAPHÚMUTATED_OUT_GRAPHÚzipÚappendr   ÚextendÚrangeÚ	recompile)ÚgmÚflat_paramsÚfw_metadataÚparamsÚfake_inp_nodesÚpreserved_arg_indicesÚout_infoÚaliased_input_argsÚiÚmÚmutated_inpsÚ
real_inputÚnodes                úV/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/_inductor/freezing.pyÚreplace_params_with_constantsr3      sE  € ð �X‰X× Ñ  MÐ Ó2€FØ˜Mœc &›kÐ*€NØÐð $×/Ñ/öàØ×ÑÐ(ð 	×ÓðÐð ô ˜k×4Ñ4Ó5÷áˆAˆqØ�?‰?Ü×)Ñ)¬<×+IÑ+IÐJñKò 	
ð€Lñ ô "+¬3¨{¸NÓ+KÓ!Lò 9ÑˆÑˆJ˜Ø�Ñ Ð%7Ñ 7Ø!×(Ñ(¨Ô+ØÜ" 2 t¨ZÕ8ð	9ð × Ñ ¤¤s¨;Ó'7¼¸V»Ó!EÔFà‡L�L„NØ Ð ùò1ùós   »D5Á48D:c                ó\   — t        «       5  t        | ||«      cddd«       S # 1 sw Y   yxY w)a5  
    Inlines parameters that are not mutated into constants and optimizes the graph through constant propagation
    and other techniques. If enabled, the function also discards the original parameters of the module for memory efficiency.

    Assumes that this function is run in dynamo tracing post aot_autograd.

    Args:
        dynamo_gm (torch.fx.GraphModule): The Dynamo constructed GraphModule.
        aot_autograd_gm (torch.fx.GraphModule): The aot_autograd constructed GraphModule to be frozen.
        example_inputs (List[torch.Tensor]): A list of example input tensors to be used in the freezing process.

    Returns:
        Tuple[torch.fx.GraphModule, List[int]]: A tuple containing the frozen GraphModule and a list of indices
        of the inputs that were preserved (not turned into constants).
    N)r   Ú_freeze)Ú	dynamo_gmÚaot_autograd_gmÚexample_inputss      r2   Úfreezer9   C   s/   € ô( 
Ó	ñ CÜ�y /°>ÓB÷C÷ Cò Cús   ‹"¢+c                ó²  — t        |«       t        j                  j                  j	                  «       x}r:|j
                  }|j                  €J ‚|j                  }|�|€J ‚t        |||«      }n9|j                  j                  d¬«      }t        t        t        |«      «      «      }t        |j                  «      }||_        |j                  «        |D �	cg c]  }	||	   ‘Œ	 }
}	t        ||
«       t!        |«       t"        j$                  rt'        «        t)        | «       t*        j-                  dt/        d|d¬«      «       t1        |«       ||fS c c}	w )Nr   r   z%szFROZEN GRAPHT)Úcolored)r   ÚtorchÚ_guardsÚTracingContextÚtry_getr'   Úparams_flat_unwrap_subclassesr3   r   r   Úlistr#   r   r	   r$   r   r
   r   Úfreezing_discard_parametersÚinvalidate_eager_modulesÚdiscard_traced_gm_paramsÚlogÚdebugr   r   )r6   r7   r8   Útracing_contextr'   Úparams_flatr*   ÚinputsÚ	cse_graphÚindÚaot_example_inputss              r2   r5   r5   [   sD  € ô �OÔ$äŸ-™-×6Ñ6×>Ñ>Ó@Ð@€Ð@Ø%×1Ñ1ˆØ×<Ñ<ÐHÐHÐHØ%×CÑCˆØÐ&¨;Ð+BÐBÐBä =Ø˜[¨+ó!
Ñð !×&Ñ&×1Ñ1°]Ð1ÓCˆÜ $¤U¬3¨v«;Ó%7Ó 8Ðô ˜_×2Ñ2Ó3€IØ%€OÔØ×ÑÔà9NÖO°#˜.¨Ó-ÐOÐÐOÜ�OÐ%7Ô8ä�/Ô"ä×)Ò)Ü Ô"Ü  Ô+ä‡I�IØÔ$ ^°_ÈdÔSôô ˜_Ô-ØÐ1Ð1Ð1ùò Ps   ÃEc                  óB   ‡ — e Zd Zeˆ fd„«       Zdd„Zedd„«       Zˆ xZS )ÚErasedTensorc                óD   •— t         ‰| �  | |j                  d¬«      «      S )NÚmeta)Údevice)ÚsuperÚ__new__Úto)ÚclsÚelemÚnameÚ
owning_modÚ	__class__s       €r2   rS   zErasedTensor.__new__‰   s   ø€ ä‰w‰˜s D§G¡G°6 GÓ$:Ó;Ð;ó    c                óF   — || _         t        j                  |«      | _        y )N)Úerased_nameÚweakrefÚrefÚowning_mod_ref)ÚselfrV   rW   Úmods       r2   Ú__init__zErasedTensor.__init__�   s   € ØˆÔÜ%Ÿk™k¨#Ó.ˆÕrZ   c           	     óò   — t        j                  |i |¤ŽD �cg c]  }t        |t        «      r|‘Œ }}t	        |«      dkD  sJ ‚|d   }t        d|› d|j                  › d|j                  «       › �«      ‚c c}w )Nr   z‰Trying to run Pytorch Eager Module after Dynamo Freezing. The original parameters have been discarded for memory efficiency. Found in op z for erased parameter z of )ÚpytreeÚarg_tree_leavesÚ
isinstancerN   r   ÚRuntimeErrorr\   r_   )rU   ÚfuncÚtypesÚargsÚkwargsÚeÚerased_tensorss          r2   Ú__torch_dispatch__zErasedTensor.__torch_dispatch__‘   s—   € ô ×+Ñ+¨TÐ<°VÑ<ö
àÜ˜!œ\Ô*ò ð
ˆð 
ô
 �>Ó" QÒ&Ð&Ð&Ø˜1Ñˆäðà˜&Ð 6°q·}±}°oÀTÈ!×JZÑJZÓJ\ÐI]ð_ó
ð 	
ùò
s   ˜A4)rW   zOptional[str]ÚreturnÚNone)© N)	Ú__name__Ú
__module__Ú__qualname__ÚstaticmethodrS   rb   Úclassmethodrn   Ú__classcell__)rY   s   @r2   rN   rN   ˆ   s.   ø„ Øó<ó ð<ó/ð ò
ó ô
rZ   rN   c            
     ó  — t         j                  j                  j                  «       5  t         j                  j
                  j                  «       j                  j                  j                  «       D ]õ  } t        | t         j                  j                  «      sŒ(t        t        j                  | j!                  d¬«      | j#                  d¬«      «      «      D ]Œ  \  }}t         j$                  j&                  j)                  «       5  t+        ||| «      }d d d «       t        |t         j                  j,                  «      rj/                  d«       d|_        t3        | |«       ŒŽ Œ÷ 	 d d d «       y # 1 sw Y   Œ`xY w# 1 sw Y   y xY w©NF)ÚrecurseT)r<   ÚutilsÚ_python_dispatchÚ_disable_current_modesr=   r>   ÚgetÚmodule_contextÚ
nn_modulesÚvaluesrf   ÚnnÚModulerA   Ú	itertoolsÚchainÚnamed_parametersÚnamed_buffersÚ	_dispatchÚpythonÚno_python_dispatcherrN   Ú	ParameterÚrequires_grad_Ú	_is_paramÚsetattr©ra   Ú	attr_nameÚtensorÚe_ts       r2   rC   rC   ¢   s8  € Ü	�‰×	%Ñ	%×	<Ñ	<Ó	>ñ -ô �]‰]×)Ñ)×-Ñ-Ó/×>Ñ>×IÑI×PÑPÓRò	-Øä˜c¤5§8¡8§?¡?Ô3Øä%)Ü—‘Ø×(Ñ(°Ð(Ó7Ø×%Ñ%¨eÐ%Ó4óó&ò -Ñ!�	˜6ô —_‘_×+Ñ+×@Ñ@ÓBñ ?Ü& v¨y¸#Ó>�C÷?ä˜f¤e§h¡h×&8Ñ&8Ô9Ø×&Ñ& tÔ,Ø$(�C”MÜ˜˜Y¨Õ,ñ-ñ	-÷-ð -÷?ð ?ú÷-ð -ús%   ©C FÄ	E6	ÄAFÅ6E?Å;FÆFc           	     ó4  — t         j                  j                  j                  «       5  t	        t        j                  | j                  d¬«      | j                  d¬«      «      «      D ]Œ  \  }}t         j                  j                  j                  «       5  t        ||| «      }d d d «       t        |t         j                  j                  «      rj!                  d«       d|_        t%        | |«       ŒŽ 	 d d d «       y # 1 sw Y   Œ^xY w# 1 sw Y   y xY wry   )r<   r{   r|   r}   rA   r„   r…   r†   r‡   rˆ   r‰   rŠ   rN   rf   r‚   r‹   rŒ   r�   rŽ   r�   s       r2   rD   rD   ¸   sè   € Ü	�‰×	%Ñ	%×	<Ñ	<Ó	>ñ )Ü!%Ü�O‰OØ×$Ñ$¨UÐ$Ó3°S×5FÑ5FÈuÐ5FÓ5Uóó"
ò 
	)ÑˆI�vô
 —‘×'Ñ'×<Ñ<Ó>ñ ;Ü" 6¨9°cÓ:�÷;ä˜&¤%§(¡(×"4Ñ"4Ô5Ø×"Ñ" 4Ô(Ø $�”Ü�C˜ CÕ(ñ
	)÷)ð )÷;ð ;ú÷)ð )ús%   ©A.DÂDÂ%ADÄDÄDÄDc                óŠ  — | j                   j                  �^ }}|j                  d   }| j                   j                  |«      5  |D ]»  }t	        |j
                  d   t        j                  «      r,t        j                  j                  |j
                  d   «      sŒW|j
                  d   }| j                   j                  t        j                  j                  ||j                  «       f«      }|j                  ||«       Œ½ 	 ddd«       | j                   j!                  «        | j#                  «        y# 1 sw Y   Œ4xY w)zì
    Make sure the output node's layout does not change due to compiler optimizations
    by adding aten.as_strided nodes with the expected strides.

    Only used for inference so we can assume all graph outputs are model outputs.
    r   ÚvalN)r   Únodesrj   Úinserting_beforerf   rP   r<   ÚTensorÚ_prims_commonÚis_non_overlapping_and_denseÚcall_functionÚprimsÚinductor_force_stride_orderÚdefaultÚstrideÚreplace_input_withÚlintr$   )r%   Ú_Úoutput_nodeÚout_listÚnÚftÚnew_nodes          r2   Úenforce_output_layoutr¨   Ç   sü   € ð —h‘h—n‘n�O€QˆØ×Ñ Ñ"€HØ	�‰×	"Ñ	" ;Ó	/ñ 8Øò 	8ˆAÜØ—‘�u‘œuŸ|™|ôä×(Ñ(×EÑEÀaÇfÁfÈUÁmÔTØð —‘˜‘ˆBØ—x‘x×-Ñ-Ü×1Ñ1×9Ñ9¸A¸r¿y¹y»{Ð;KóˆHð ×*Ñ*¨1¨hÕ7ñ	8÷8ð$ ‡H�H‡M�M„OØ‡L�L…N÷'8ð 8ús   ÁCD9Ä9Ec                ó^  — t         j                  j                  j                  j                  t         j                  j                  j
                  j                  t         j                  j                  j                  j                  g}| j                  j                  D �cg c]  }|j                  |v sŒ|‘Œ }}|D ]²  }| j                  j                  |«      5  |j                  d   j                  d   }| j                  j                  t        j                  j                  |j                  d   |j!                  «       f«      }|j#                  |j                  d   |«       ddd«       Œ´ | j                  j%                  «        | j'                  «        yc c}w # 1 sw Y   ŒîxY w)z´
    Make sure the as_strided node's input's layout does not change due to compiler
    optimizations, because the as_strided strides info depends on input tensor stride info.
    r   r•   N)r<   ÚopsÚatenÚ
as_stridedrž   Úas_strided_Úas_strided_scatterr   r–   Útargetr—   rj   rP   r›   rœ   r�   rŸ   r    r¡   r$   )r%   Úas_strided_opsr¥   Ústrided_nodesr¦   r§   s         r2   Úenforce_as_strided_input_layoutr²   æ   s9  € ô 	�	‰	�‰×!Ñ!×)Ñ)Ü�	‰	�‰×"Ñ"×*Ñ*Ü�	‰	�‰×)Ñ)×1Ñ1ð€Nð
 !#§¡§¡ÖM˜1°!·(±(¸nÒ2L’QÐM€MÐMØò 6ˆØ�X‰X×&Ñ& qÓ)ñ 	6à—‘˜‘—‘ Ñ&ˆBØ—x‘x×-Ñ-Ü×1Ñ1×9Ñ9¸A¿F¹FÀ1¹IÀrÇyÁyÃ{Ð;SóˆHð × Ñ  §¡¨¡¨HÔ5÷	6ð 	6ð6ð ‡H�H‡M�M„OØ‡L�L…Nùò N÷	6ð 	6ús   Â"FÂ6FÃBF#Æ#F,	c           	     óü  — t        d«      5  | j                  j                  D �cg c],  }|j                  t        j
                  j                  k(  sŒ+|‘Œ. }}|D ]ä  }|j                  d   }t        |j                  d   j                  «       «      dk7  s-|j                  d   j                  t        j                  ¬«      rŒi| j                  j                  |«      5  | j                  j                  t        j                   j                  |fdt        j                  i«      }|j#                  ||«       ddd«       Œæ t%        | «       t'        | «       ddd«       yc c}w # 1 sw Y   �ŒxY w# 1 sw Y   yxY w)z®
    Convert 4d convolution weight tensor to channels last format.

    This pass is performed before freezing so the added nodes can be constant
    folded by freezing.
    Ú%convert_conv_weights_to_channels_lastr   r•   é   )Úmemory_formatr¶   N)r   r   r–   r¯   r«   Úconvolutionrž   rj   r   rP   ÚsizeÚis_contiguousr<   Úchannels_lastr—   r›   Úcloner    r²   r¨   )r%   r¥   ÚconvsÚconvÚweight_noder§   s         r2   r´   r´   ÿ   sI  € ô 
Ð=Ó	>ñ "ØŸH™HŸN™NÖS�q¨a¯h©h¼$×:JÑ:J×:RÑ:RÓ.R’ÐSˆÐSØò 	?ˆDØŸ)™) A™,ˆKÜ�;×#Ñ# EÑ*×/Ñ/Ó1Ó2°aÒ7¸;×;KÑ;KØñ<ç‰m¬%×*=Ñ*=ˆmÓ>ð<?ð à—‘×*Ñ*¨4Ó0ñ ?ØŸ8™8×1Ñ1Ü—J‘J×&Ñ&Ø �NØ$¤e×&9Ñ&9Ð:ó�ð
 ×'Ñ'¨°XÔ>÷?ð ?ð	?ô  	(¨Ô+Ü˜bÔ!÷'"ð "ùÚS÷?ñ ?ú÷"ð "ús<   ŒE2¥,E ÁE ÁBE2ÃAE%Ä7 E2Å E2Å%E/Å*E2Å2E;)r%   útorch.fx.GraphModuler&   z	list[Any]r'   z1torch._functorch.aot_autograd.ViewAndMutationMetaro   z	list[int])r6   r¿   r7   r¿   r8   z"list[torch._subclasses.FakeTensor]ro   z&tuple[torch.fx.GraphModule, list[int]])ra   r¿   )r%   r¿   )0Ú
__future__r   r„   Úloggingr]   Útypingr   r   r<   Útorch.utils._pytreer{   Ú_pytreerd   Útorch._dynamo.utilsr   r   Útorch._functorch.aot_autogradr   Útorch._functorch.compile_utilsr	   Ú torch._inductor.constant_foldingr
   r   Útorch._inductor.freezing_utilsr   r   Ú+torch._inductor.fx_passes.freezing_patternsr   Ú#torch._inductor.fx_passes.post_gradr   Ú r   rª   r«   rœ   Ú	getLoggerrr   rE   r3   r9   r5   r˜   rN   rC   rD   r¨   r²   r´   rq   rZ   r2   ú<module>rÎ      s  ðå "ã Û Û ß  ã ß $Ð $ß DÝ 6Ý 7ß Vß SÝ GÝ ?å ð ‡y�y‡~�~€Ø�	‰	�‰€à€g×Ñ˜Ó!€ð$!Øð$!àð$!ð Cð$!ð ó	$!ðNCØ#ðCà)ðCð 7ðCð ,ó	Cð0*2Ø#ð*2à)ð*2ð 7ð*2ð ,ó	*2ôZ
�5—<‘<ô 
ò4-ó,)óó>ô2"rZ   