ó
    EñiJ4  ã                   ó4  • S SK r S SKJr  S SKJr  S SKrS SKrS SKJr  S SK	J
r
  S SKJrJrJr  S SKJr  S SKJr  S S	KJr  S S
KJr  S SKJrJrJr  S SKJr  S SKJr   " S S\5      r \ " 5       r!\!RE                  \5      S 5       r#\!RE                  \5      S 5       r$\!RJ                  S 5       r&\!RE                  \RN                  5      " \" \!SS95        \!RE                  \RP                  5      S 5       r)S r*\S 5       r+S r,S r-S\.S\.4S jr/S r0S r1S r2g) é    N)Úcontextmanager)Úwraps)ÚDispatchKey)Ú+_maybe_find_pre_dispatch_tf_mode_for_export)Ú_ConstantFunctionÚ
flat_applyÚto_graphable©Ústrict_mode)Úautograd_not_implemented)ÚHigherOrderOperator)ÚFakeTensorMode)ÚPreDispatchTorchFunctionModeÚProxyTorchDispatchModeÚtrack_tensor_tree)Ú_pytree)Ú"is_traceable_wrapper_subclass_typec                   ó4   ^ • \ rS rSrU 4S jrU 4S jrSrU =r$ )ÚExportTracepointé   c                 ó$   >• [         TU ]  S5        g )NÚ_export_tracepoint)ÚsuperÚ__init__)ÚselfÚ	__class__s    €ÚS/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/_export/wrappers.pyr   ÚExportTracepoint.__init__   s   ø€ Ü‰ÑÐ-Õ.ó    c                 ó$   >• [         TU ]  " U0 UD6$ ©N)r   Ú__call__)r   ÚargsÚkwargsr   s      €r   r"   ÚExportTracepoint.__call__    s   ø€ ä‰wÒ Ð0¨Ñ0Ð0r   © )Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__r   r"   Ú__static_attributes__Ú__classcell__)r   s   @r   r   r      s   ø† õ/÷1ó 1r   r   c                 óÊ   • [         R                  " U R                  R                  X45      u  p4U R                  R	                  S[
        X45      n[        XS U R                  S9$ )NÚcall_function©ÚconstantÚtracer)ÚpytreeÚtree_mapr1   Úunwrap_proxyÚcreate_proxyr   r   )Úmoder#   r$   Úp_argsÚp_kwargsÚproxys         r   Úexport_tracepoint_dispatch_moder:   (   sR   € ä—’ t§{¡{×'?Ñ'?À$ÀÓPÑ€FØ�K‰K×$Ñ$ØÔ+¨Vó€Eô ˜T°4ÀÇÁÑLÐLr   c                 ó@   • U    UsS S S 5        $ ! , (       d  f       g = fr!   r&   )r6   r#   r$   s      r   Ú"export_tracepoint_fake_tensor_moder<   1   s   € â	Ø÷ 
��ús   ƒ�
c                 ó¶   • U R                  U5      nU R                  U5      nU R                  5          [        U0 UD6  UsS S S 5        $ ! , (       d  f       g = fr!   )Úunwrap_tensorsÚredispatch_to_nextr   )Úctxr#   r$   Úunwrapped_argsÚunwrapped_kwargss        r   Úexport_tracepoint_functionalrC   7   sL   € à×'Ñ'¨Ó-€NØ×)Ñ)¨&Ó1Ðà	×	Ñ	Õ	!Ü˜NÐ?Ð.>Ò?Ø÷ 
"×	!×	!ús   ³A
Á

AT)Údeferred_errorc                  ó   • U $ r!   r&   )r#   r$   s     r   Úexport_tracepoint_cpurF   F   s   € à€Kr   c                 óv  ^^^^	• [        U [        R                  R                  5      (       d  [	        S[        U 5       35      eTS:X  a  [	        S5      e[        R                  R                  R                  U T5      nU4S jm	S mUU4S jnUUU	4S jnUR                  USS	9nUR                  USS	9nXg4$ )
Nzexpected torch.nn.Module, got Ú zpath must not be emptyc                 ó¸   >• U T;   aL  TU    S   U:w  a  [        SU  STU    S    SU 35      eTU    S   U:w  a  [        SU  STU    S    SU 35      eXS.TU '   g )NÚin_speczin_spec mismatch for z: z != Úout_speczout_spec mismatch for )rJ   rK   )ÚAssertionError)ÚpathrJ   rK   Úmodule_call_specss      €r   Úupdate_module_call_signaturesÚ6_wrap_submodule.<locals>.update_module_call_signaturesR   s¦   ø€ ØÐ$Ó$Ø  Ñ& yÑ1°WÓ<Ü$Ø+¨D¨6°Ð4EÀdÑ4KÈIÑ4VÐ3WÐW[Ð\cÐ[dÐeóð ð ! Ñ& zÑ2°hÓ>Ü$Ø,¨T¨F°"Ð5FÀtÑ5LÈZÑ5XÐ4YÐY]Ð^fÐ]gÐhóð ð /6Ñ"LÐ˜$Òr   c           	      ó¤   • U  HJ  n[        U[        R                  [        [        [
        [        45      (       a  M9  Uc  M>  [        SU 35      e   g )NzGOnly Tensors or scalars are supported as pytree flattened inputs, got: )Ú
isinstanceÚtorchÚTensorÚstrÚintÚfloatÚboolrL   )Ú	flat_argsÚas     r   Úcheck_flattenedÚ(_wrap_submodule.<locals>.check_flattened^   sE   € ÛˆAÜ˜q¤5§<¡<´´c¼5Ä$Ð"G×HÓHÈAËIÜ$Ø]Ð^_Ð]`Ðaóð ò r   c                 ó”   >• [         R                  " X45      u  p4T" U5        [        USTS.6n[         R                  " X45      u  pX4$ )NÚmodule_call_inputs©ÚkindrM   ©r2   Útree_flattenr   Útree_unflatten)Úmoduler#   r$   rY   rJ   r[   rM   s        €€r   Úpre_hookÚ!_wrap_submodule.<locals>.pre_hooke   sJ   ø€ Ü#×0Ò0°$°Ó@Ñˆ	Ù˜	Ô"Ü&¨	Ð8LÐSWÒXˆ	Ü×,Ò,¨YÓ@‰ˆØˆ|Ðr   c                 óÌ   >• [         R                  " X45      u  pE[         R                  " U5      u  pgT" U5        [        UST	S.6nT
" T	XW5        [         R                  " Xg5      $ )NÚmodule_call_outputsr_   ra   )rd   r#   r$   ÚresÚ_rJ   Úflat_resrK   r[   rM   rO   s           €€€r   Ú	post_hookÚ"_wrap_submodule.<locals>.post_hookl   s]   ø€ Ü×(Ò(¨$¨Ó8‰
ˆÜ#×0Ò0°Ó5ÑˆÙ˜Ô!Ü% xÐ6KÐRVÒWˆÙ% d¨GÔ>Ü×$Ò$ XÓ8Ð8r   T)Úwith_kwargs)rR   rS   ÚnnÚModulerL   ÚtypeÚfxÚgraph_moduleÚ	_get_attrÚregister_forward_pre_hookÚregister_forward_hook)
ÚmodrM   rN   Ú	submodulere   rl   Ú
pre_handleÚpost_handler[   rO   s
    ``     @@r   Ú_wrap_submoduler{   K   s¨   û€ Ü�cœ5Ÿ8™8Ÿ?™?×+Ñ+ÜÐ=¼dÀ3»i¸[ÐIÓJÐJØˆrƒzÜÐ5Ó6Ð6Ü—‘×%Ñ%×/Ñ/°°TÓ:€Iõ
Mòö÷9ð ×4Ñ4°XÈ4Ð4ÐP€JØ×1Ñ1°)ÈÐ1ÐN€KØÐ"Ð"r   c              #   óÐ   #   • / n U H  nUR                  [        XU5      5        M      S v •  U H  nUR                  5         M     g ! U H  nUR                  5         M     f = f7fr!   )Úextendr{   Úremove)ÚfÚpreserve_signatureÚmodule_call_signaturesÚhandlesrM   Úhandles         r   Ú_wrap_submodulesr„   y   sV   é € à€GðÛ&ˆDØ�N‰Nœ?¨1Ð4JÓKÖLñ 'ããˆFØ�M‰MŽOò ø“gˆFØ�M‰MŽOò üs   ‚A&†(A ®A&ÁA#Á#A&c                 ó   • S nXl         U $ )Nc                 ó   • [        X5      $ r!   r
   )r   r#   s     r   ÚcallÚ'_mark_strict_experimental.<locals>.call‡   s   € Ü˜4Ó&Ð&r   )r"   )Úclsr‡   s     r   Ú_mark_strict_experimentalrŠ   †   s   € ò'ð „LØ€Jr   c                 ó0  • US-   n[        U R                  U5      (       a<  [        U R                  U5      U:w  a  [        SU 35      eU R	                  SUS0 5      $ U R                  U5      n[        U R                  XB5        U R	                  SUS0 5      $ )a  
This is a wrapper utility method on top of tracer to cache the
already registered subclass spec attribute. This is useful because
Subclass.__init__ will be same for each subclass. By default, fx will
create multiple attributes/proxies for given attribute.
Ú0zspec mismatch for Úget_attrr&   )ÚhasattrÚrootÚgetattrrL   r5   Úget_fresh_qualnameÚsetattr)r1   ÚnameÚspecÚfx_nameÚqualnames        r   Ú#_register_func_spec_proxy_in_tracerr—   Ž   s�   € ð �S‰j€GÜˆv�{‰{˜G×$Ñ$Ü�6—;‘; Ó(¨DÓ0Ü Ð#5°g°YÐ!?Ó@Ð@Ø×"Ñ" :¨w¸¸BÓ?Ð?à×(Ñ(¨Ó.€HÜˆF�K‰K˜Ô(Ø×Ñ˜z¨8°R¸Ó<Ð<r   Ú	spec_nameÚcall_spec_cache_keyc                 ó„  • [        U5      u  pgU R                  U5      n[        U R                  X‡5        U R	                  SUS0 5      n	[
        R                  " [        U5      5      u  p«[        X S3U5      n[
        R                  " U R                  U5      nU R	                  S[        XÉ/UQ70 5      n[        XNS U S9  g )Nr�   r&   Ú_const_func_specr.   r/   )r	   r‘   r’   r�   r5   r2   rb   r   r—   r3   r4   r   r   )r1   r˜   Úconst_target_for_applyÚgraphable_argsÚtrack_valuer™   rY   rJ   r–   Ú
spec_proxyrj   Ú	func_specÚfunc_spec_proxyÚflat_proxy_argsÚ	out_proxys                  r   Ú_emit_flat_apply_callr¤       sÂ   € ô & nÓ5Ñ€IØ×(Ñ(¨Ó3€HÜˆF�K‰K˜Ô+Ø×$Ñ$ Z°¸2¸rÓB€Jô ×&Ò&Ô'8Ð9OÓ'PÓQ�L€AÜ9ØÐ'Ð'7Ð8¸)ó€Oô
 —o’o f×&9Ñ&9¸9ÓE€Oð ×#Ñ#Øœ oÐ%TÀOÑ%TÐVXó€Iô �k°tÀFÓKr   c                 óD   • [        U 5      =(       a    U R                  S:H  $ )Nr   )Úcallabler'   )Úfns    r   Ú_is_initr¨   ¿   s   € Ü�B‹<×5˜BŸK™K¨:Ñ5Ð5r   c                 óf   ^ • [        T 5      (       d  [        ST R                   S35      eU 4S jnU$ )aé  
Experimental decorator that makes subclass to be traceable in export
with pre-dispatch IR. To make your subclass traceble in export, you need to:
    1. Implement __init__ method for your subclass (Look at DTensor implementation)
    2. Decorate your __init__ method with _mark_constructor_exportable_experimental
    3. Put torch._dynamo_disable decorator to prevent dynamo from peeking into its' impl

Example:

class FooTensor(torch.Tensor):
    @staticmethod
    def __new__(cls, elem, *, requires_grad=False):
        # ...
        return torch.Tensor._make_subclass(cls, elem, requires_grad=requires_grad)

    @torch._dynamo_disable
    @mark_subclass_constructor_exportable_experimental
    def __init__(self, elem, ...):
        # ...
z‰torch._export.wrappers.mark_constructor_exportable_experimental can only be applied on subclass tensor.__init__But, you are adding it on z‡ which is not supported. If __init__ doesn't exist on your subclass, please add it. Look at DTensor.__init__ implementation for examplec            	      óð  >• T	" U 0 UD6  [         R                  R                  5       (       d  g [        [	        U S   5      5      (       d`  T	R
                  R                  S5      (       d  [        ST	R
                   35      eT	R
                  S [        S5      *  n[        SU S35      e[        5       nUc  g [        U[        5      (       d  [        S[	        U5       35      eUR                  nU S   n[        U SS  5      U4nSR                  T	R
                  R!                  5       R#                  S	5      5      n[	        U5      R$                  R!                  5       n['        UU[	        U5      UUUS
9  g )Nr   r   z2expected __qualname__ to end with '__init__', got zCan't intercept zœ in export because this object is not a traceable tensor subclass. Please look at DTensor.__init__ implementation as an example of proper usage of this API.z+expected PreDispatchTorchFunctionMode, got é   rj   Ú.©r1   r˜   rœ   r�   rž   r™   )rS   ÚcompilerÚis_exportingr   rq   r)   ÚendswithrL   ÚlenÚRuntimeErrorr   rR   r   r1   ÚtupleÚjoinÚlowerÚsplitr'   r¤   )
r#   r$   Úobj_namer6   r1   ÚsubclassÚ	graphabler˜   r™   Úconstructor_subclasss
            €r   ÚwrapperÚBmark_subclass_constructor_exportable_experimental.<locals>.wrapperß   ss  ø€ Ù˜dÐ- fÒ-ä�~‰~×*Ñ*×,Ñ,Øä1´$°t¸A±w³-×@Ñ@Ø'×4Ñ4×=Ñ=¸j×IÑIÜ$ØHÐI]×IjÑIjÐHkÐlóð ð ,×8Ñ8Ð9K¼CÀ
»OÐ;KÐLˆHÜØ" 8 *ð -}ð ~óð ô
 ;Ó<ˆØ‰<Øä˜$Ô <×=Ñ=Ü Ø=¼dÀ4»j¸\ÐJóð ð —‘ˆØ˜‘7ˆÜ˜4  ˜8“_ fÐ-ˆ	à—H‘HÐ1×>Ñ>×DÑDÓF×LÑLÈSÓQÓRˆ	Ü" 8›n×5Ñ5×;Ñ;Ó=ÐäØØÜ#'¨£>Ø$Ø Ø 3ò	
ð 	r   )r¨   r²   r'   )rº   r»   s   ` r   Ú1mark_subclass_constructor_exportable_experimentalr½   Ã   sI   ø€ ô* Ð(×)Ñ)Üð)Ø)=×)FÑ)FÐ(Gð H}ð~ó
ð 	
õ)ðV €Nr   c                 óØ   ^ • [        T 5      (       a  [        T 5      $ [        T 5      (       d)  T R                  S:X  d  [        ST R                   S35      e[	        T 5      U 4S j5       nU$ )ao  
Experimental decorator that adds user function to export pre-dispatch graph. Note that
we only support custom autograd function/subclass constructors today. To use this function:
    1. For subclasses:
        1. refer to instructions in mark_subclass_constructor_exportable_experimental
    2. Define apply method on your custom autograd function and apply this decorator.

Example:

class MyCoolCustomAutogradFunc(autograd.Function):
    @classmethod
    @torch._export.wrappers.allow_in_pre_dispatch_graph
    def apply(cls, *args, **kwargs):
        return super(MyCoolCustomAutogradFunc, cls).apply(*args, **kwargs)

ÚapplyzŸtorch._export.wrappers.allow_in_pre_dispatch_graph can only be applied on subclass tensor.__init_ or custom_autograd_function.apply. But, you are adding it on a/   which is not supported. If __init__ doesn't exist on your subclass, please add it. Look at DTensor.__init__ implementation for example. If you are adding it on custom autograd function, please add it on apply method. If anything else, file an issue on github and we may consider extending our support. c            	      ó¾  >• [         R                  R                  5       (       d  T" U 0 UD6$ [        R                  " U S   5      (       d  T" U 0 UD6$ [        U S   [         R                  R                  5      (       d  T" U 0 UD6$ SSKJ	n  U" [         R                  R                  R                  5      nUc  T" U 0 UD6$ [         R                  R                  5       R                  [         R                  R                  R                   5      n[         R                  R#                  5       [         R                  R%                  [         R                  R                  R                   5      -  n[         R                  R'                  XE5         T" U 0 UD6nS S S 5        UR(                  (       d  [+        S5      eUR,                  nU S   R.                   SU S   R0                   3nU/U SS  Q7U4n	SSKJn
  SR7                  UR9                  S5      5      n[;        U
5      R<                  R?                  5       n[A        UUU
U	WUS9  U$ ! , (       d  f       N»= f)	Nr   )Ú_get_dispatch_mode_pre_dispatchz"Should only do this in predispatchr¬   r«   )Ú._call_custom_autograd_function_in_pre_dispatchrj   r­   )!rS   r®   r¯   ÚinspectÚisclassÚ
issubclassÚautogradÚFunctionÚ
torch._opsrÁ   Ú_CÚ_TorchDispatchModeKeyÚPROXYÚ_dispatch_tls_local_include_setr~   r   ÚPreDispatchÚ_dispatch_tls_local_exclude_setÚDispatchKeySetÚ_ForceDispatchKeyGuardÚpre_dispatchrL   r1   r(   r)   Útorch.export.custom_opsrÂ   r´   r¶   rq   r'   rµ   r¤   )r#   r$   rÁ   r6   Úinclude_to_setÚexclude_to_setÚoutr1   Úfunction_cls_namer¹   rÂ   r˜   r™   Úfuncs                €r   r»   Ú,allow_in_pre_dispatch_graph.<locals>.wrapper+  s  ø€ ä�~‰~×*Ñ*×,Ñ,Ù˜Ð( Ñ(Ð(ä�Š˜t A™w×'Ñ'Ù˜Ð( Ñ(Ð(ä˜$˜q™'¤5§>¡>×#:Ñ#:×;Ñ;Ù˜Ð( Ñ(Ð(å>á.¬u¯x©x×/MÑ/M×/SÑ/SÓTˆØ‰<Ù˜Ð( Ñ(Ð(ô Ÿ™×AÑAÓC×JÑJÜ�H‰H× Ñ ×,Ñ,ó
ˆô �H‰H×4Ñ4Ó6Ü�h‰h×%Ñ%¤e§h¡h×&:Ñ&:×&FÑ&FÓGñHð 	ô
 �X‰X×,Ñ,¨^ÕLÙ˜Ð' Ñ'ˆC÷ Mð × × Ü Ð!EÓFÐFØ—‘ˆà# A™w×1Ñ1Ð2°!°D¸±G×4HÑ4HÐ3IÐJÐØ'Ð3¨$¨q¨r¨(Ñ3°VÐ<ˆ	õ	
ð —H‘HÐ.×4Ñ4°SÓ9Ó:ˆ	Ü"Ø:ó
ç
‰(—5‘5“7ð 	ô 	ØØØ#QØ$ØØ 3ò	
ð ˆ
÷5 MÕLús   Æ	IÉ
I)r¨   r½   r'   r²   r   )r×   r»   s   ` r   Úallow_in_pre_dispatch_graphrÙ     st   ø€ ô" �‡~�~Ü@ÀÓFÐFä�T�N‰N˜dŸm™m¨wÓ6Üð)à)-¯©¨ð 8dðeó
ð 	
ô ˆ4ƒ[ô4ó ð4ðl €Nr   )3rÃ   Ú
contextlibr   Ú	functoolsr   rS   Útorch._custom_opsÚtorch._Cr   Útorch._export.utilsr   Ú"torch._higher_order_ops.flat_applyr   r   r	   Ú#torch._higher_order_ops.strict_moder   Útorch._higher_order_ops.utilsr   rÈ   r   Útorch._subclasses.fake_tensorr   Ú"torch.fx.experimental.proxy_tensorr   r   r   Útorch.utilsr   r2   Útorch.utils._python_dispatchr   r   r   Úpy_implr:   r<   Úpy_functionalize_implrC   ÚAutogradÚCPUrF   r{   r„   rŠ   r—   rU   r¤   r¨   r½   rÙ   r&   r   r   Ú<module>rê      sR  ðã Ý %Ý ã Û Ý  Ý K÷ñ õ
 <Ý BÝ *Ý 8÷ñ õ
 *Ý Kô1Ð*ô 1ñ &Ó'Ð ð ×ÑÐ2Ó3ñMó 4ðMð ×Ñ˜NÓ+ñó ,ðð
 ×)Ñ)ñó *ðð × Ñ ˜;×/Ñ/Ô 0ÙÐ/ÀÑEôð
 ×Ñ˜KŸO™OÓ,ñó -ðò+#ð\ ñ	ó ð	òò=ð$Lð ðLð ôLò>6òGóTUr   