
    ^j                       d dl mZ d dlZd dlZd dlZd dlZd dlZd dlZd dlZd dl	Z	d dl
Z
d dlZd dlZd dlmZ d dlmZ d dlmZ d dlmZ d dlmZmZmZmZmZmZ d dlmZ d dlZd d	lmZ  d d
l!m"Z" d dl#m$Z$ d dl%m&Z&m'Z' d dl(m)Z)  ejT                  e+      Z,erBd dl-Z-d dl.m/Z/m0Z0m1Z1 d dl2m3Z3 d dl4Z4d dl5m6Z6 d dl7m8Z8 d dl9m:Z: d dl;m<Z< d dl=m>Z> d dl?m@Z@ d dlAmBZB 	  ej                  d      ZD ej                  d      ZE eddd       G d d             ZF G d de      ZG G d d ej                        ZI	  G d! d"      ZJ ej&                  d#       G d$ d%             ZK G d& d'e      ZL ej&                  d(       G d) d*             ZM ed+      ZN	  ej&                  d#       G d, d-             ZO	  ej&                  d#       G d. d/eO             ZP	  ej&                  d#       G d0 d1eO             ZQ	  G d2 d3eeN         ZR G d4 d5      ZS G d6 d7      ZT G d8 d9eReT         ZU G d: d;      ZV G d< d=eReV         ZW G d> d?      ZX	  G d@ dAeReS         ZY G dB dC      ZZe G dD dE             Z[e G dF dG             Z\ G dH dIeZ      Z] G dJ dK      Z^ e	j                         Z`	  G dL dM      Zae G dN dO             Zb G dP dQ      Zce	 	 	 	 d]dR       Zde	 	 	 	 d^dS       Zeed_dT       Zfe	 d`	 	 	 	 	 dadU       Zf e       	 d`	 	 	 	 	 dbdV       Zf efd#       G dW dX             Zg efd#       G dY dZeg             Zhd`dcd[Ziddd\Zjy)e    )annotationsN)abstractmethod)defaultdict)contextmanager)	dataclass)AnyGeneric
NamedTupleoverloadTYPE_CHECKINGTypeVar)dataclass_transform)_pytree)
OrderedSet)is_traceable_wrapper_subclass)CapturedTracebackformat_frame)WeakTensorKeyDictionary)Callable	GeneratorIterator)CodeType)DDPOptimizerContext)	PyCodegen)GuardCheckSpec)CodeOptions)ViewAndMutationMeta)NestedCompileRegionOptionsFakeTensorModez-^(?P<frame_id>\d+)/(?P<frame_compile_id>\d+)$zQ^!(?P<compiled_autograd_id>\d+)(?:/(?P<frame_id>\d+)/(?P<frame_compile_id>\d+))?$T)frozenkw_onlyslotsc                  J    e Zd ZU ded<   ded<   dZded<   ddZed	d       Zy)
	CompileId
int | Noneframe_idframe_compile_idNcompiled_autograd_idc                   | j                   }| j                  d u | j                  d u k7  r%t        d| j                   d| j                         d}| j                  d| j                   d| j                   }d| j                    | S | j                  | j                  %t        d| j                   d| j                         | j                   d| j                   S )NzMframe_id and frame_compile_id must both be None or both be set, got frame_id=z, frame_compile_id= /!z=frame_id and frame_compile_id must not be None, got frame_id=)r)   r'   r(   AssertionError)self	frame_strs     X/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/torch/_guards.py__str__zCompileId.__str__X   s   $$0%4+@+@D+HI$$$(MM?2EdF[F[E\^  I}}(a0E0E/FG	t001)==}}$(=(=(E$$$(MM?2EdF[F[E\^  mm_Ad&;&;%<==    c                &   |y	 t         t        fD ]X  }|j                  |      x}s|j                         }|j	                         D ]  \  }}|	t        |      ||<     | di |c S  t        # t        $ r}t        d| d      |d}~ww xY w)z
        Factory method that creates a CompileId from its string representation.
        Keep this in sync with the __str__ method.
        NzInvalid compile_id '' )COMPILE_ID_PATTERNCA_COMPILE_ID_PATTERNmatch	groupdictitemsint
ValueError	Exception)cls
compile_idpatternr9   groupskves           r1   from_stringzCompileId.from_stringm   s     	J.0EF !#MM*5555"__.F & /1=(+AF1I/ ==(! !  	J3J<qABI	Js(   "A2 (A2 A2 +A2 2	B;BBreturnstr)r@   
str | NonerH   CompileId | None)__name__
__module____qualname____annotations__r)   r2   classmethodrF   r6   r3   r1   r%   r%   I   s8    
 !  (,*+
>* J Jr3   r%   c                  *    e Zd ZU ded<   ded<   ddZy)TraceIdr%   r@   r<   attemptc                    | j                   dk(  rt        | j                        S | j                   d| j                    S )Nr   _)rS   rI   r@   r/   s    r1   r2   zTraceId.__str__   s7    <<1t''oo&a~66r3   NrG   rL   rM   rN   rO   r2   r6   r3   r1   rR   rR      s     L7r3   rR   c                  x    e Zd ZdZdZdZdZdZdZdZ	dZ
d	Zd
ZdZdZdZdZdZdZdZddZddZddZddZddZy)GuardSourcer                           	   
                     c                F    | t         j                  t         j                  fv S N)rY   GLOBAL_FSDP_MODULELOCAL_FSDP_MODULErV   s    r1   is_fsdp_modulezGuardSource.is_fsdp_module   s    668U8UVVVr3   c                    dd l mc m} |j                  r4| t        j
                  t        j                  fv xs | j                         S | t        j
                  t        j                  fv S Nr   )torch._dynamo.config_dynamoconfig_unsafe_skip_fsdp_module_guardsrY   GLOBAL_SPECIALIZED_NN_MODULELOCAL_SPECIALIZED_NN_MODULErn   )r/   rs   s     r1   is_specialized_nn_modulez$GuardSource.is_specialized_nn_module   sn    --11<<;; )
 &&( 4433
 
 	
r3   c                    | t         j                  t         j                  t         j                  t         j                  fv S rk   )rY   GLOBAL_UNSPECIALIZED_NN_MODULELOCAL_UNSPECIALIZED_NN_MODULE&GLOBAL_UNSPECIALIZED_BUILTIN_NN_MODULE%LOCAL_UNSPECIALIZED_BUILTIN_NN_MODULErV   s    r1   is_unspecialized_nn_modulez&GuardSource.is_unspecialized_nn_module   s8    6655>>==	
 
 	
r3   c                F    | t         j                  t         j                  fv S rk   )rY   r{   r|   rV   s    r1   "is_unspecialized_builtin_nn_modulez.GuardSource.is_unspecialized_builtin_nn_module   s&    >>==
 
 	
r3   c                    | t         j                  t         j                  t         j                  t         j                  t         j
                  fv S rk   )rY   LOCALrv   rm   rz   r|   rV   s    r1   is_localzGuardSource.is_local   sA    33))55==
 
 	
r3   NrH   bool)rL   rM   rN   r   GLOBALrv   ru   CONSTANTRANDOM_VALUE	SHAPE_ENVrm   rl   BACKWARD_STATE	EPHEMERALSYNTHETIC_LOCALrz   ry   r|   r{   
TEMP_LOCALrn   rw   r}   r   r   r6   r3   r1   rY   rY      sy    EF"##$ HLINIO$&!%'",.)-/*JW
"


r3   rY   c                      e Zd Zy)GuardBuilderBaseNrL   rM   rN   r6   r3   r1   r   r      s    r3   r   )r!   c                  *    e Zd ZU ded<   ded<   ddZy)SLocz#traceback.FrameSummary | str | Noneframework_locrJ   maybe_user_locc                    t        | j                  t              r| j                  n"| j                  dnt        | j                        }| j                  | j                   d| dS d| dS )Nr+   z ()()
isinstancer   rI   r   r   )r/   flocs     r1   r2   zSLoc.__str__   s{     $,,c2  !!) d001 	 *))*"TF!44tfA;r3   NrG   rW   r6   r3   r1   r   r      s    66r3   r   c                  ,    e Zd ZU ded<   ded<   ded<   y)
ShapeGuardzsympy.logic.boolalg.Booleanexprr   slocr   size_obliviousNrL   rM   rN   rO   r6   r3   r1   r   r      s    
%%
Jr3   r   )r#   c                  <   e Zd ZU ded<   ded<   dZded<   dZded<   dZd	ed
<   dZded<   dZded<   dZ	ded<   dZ
ded<   dZded<   d%dZd&dZd'dZd(dZed)d       Zed*d       Zed+d       Zd)dZd)dZd,dZd-d Zd-d!Zd-d"Zd)d#Z	 	 	 	 	 	 	 	 	 	 d.d$Zy)/GuardSourceoriginating_sourcez)Callable[[GuardBuilderBase, Guard], None]	create_fnNzlist[str] | Noneguard_types	code_listzobject | Noneobj_weakref!weakref.ReferenceType[Any] | Noneguarded_class_weakrefzCapturedTraceback | Nonestackztraceback.StackSummary | None
user_stackr&   _hashFr   _unserializablec                    | j                   :t        | j                  | j                  t	        | j
                        f      | _         | j                   S rk   )r   hashnamesourceidr   rV   s    r1   __hash__zGuard.__hash__  s:    ::tyy$++r$..7IJKDJzzr3   c                   t        | j                  t        j                        xrD | j                  j                  t
        j                  j                  j                  j                  u }|| j                  r| j                  j                  ndt        | j                        | j                  | j                         j                  j                   fS )N)r   r   	functoolspartialfunctorchrr   guardsGuardBuilderDUPLICATE_INPUTr   valuelenr   inner_create_fn__code__co_firstlineno)r/   is_duplicate_inputs     r1   sort_keyzGuard.sort_key"  s    
 t~~y'8'89 Y##u}}';';'H'H'X'XX 	
 !%DKK"		NII  "++::
 	
r3   c                D    | j                         |j                         k  S rk   )r   r/   others     r1   __lt__zGuard.__lt__2  s    }}!111r3   c                    t        | j                  t        j                        r| j                  j                  S | j                  S rk   )r   r   r   r   r   rV   s    r1   r   zGuard.inner_create_fn5  s0    dnni&7&78>>&&&>>!r3   c                .    | j                   j                  S rk   )r   r   rV   s    r1   r   z
Guard.name;  s    &&+++r3   c                .    | j                   j                  S rk   )r   guard_sourcerV   s    r1   r   zGuard.source?  s    &&333r3   c           
        t        | t        j                        re |        }|Ddt        t	        |              d|j
                  j                   dt        t	        |             dS dt        t	        |              dS t        |       S )a  
        This is a workaround of a Python weakref bug.

        `obj_weakref` is instance returned by `weakref.ref`,
        `str(obj_weakref)` is buggy if the original obj overrides __getattr__, e.g:

            class MyConfig(dict):
                def __getattr__(self, x):
                    return self[x]

            obj = MyConfig(offset=5)
            obj_weakref = weakref.ref(obj)
            str(obj_weakref)  # raise error: KeyError: '__name__'
        z<weakref at z; to 'z' at >z; dead>)r   weakrefReferenceTypehexr   	__class__rL   rI   )r   objs     r1   weakref_to_strzGuard.weakref_to_strC  s      k7#8#89-C%c"[/&:%;6#--BXBXAYY^_bceficj_k^llmnn%c"[/&:%;7CC{##r3   c                Z   d| j                   r$| j                   j                  j                         nd dt        | j                         d| j	                         j
                   d| j                   d| j                   d| j                  | j                         d| j                   d}|S )	Nz	
        r+    z&
        {
            'guard_types': z,
            'code': z,
            'obj_weakref': z
            'guarded_class': z
        }
        )r   r   lowerreprr   rL   r   r   r   r   r   )r/   ss     r1   __repr__zGuard.__repr__\  s    	%)[[					!b94		?:K1TMaMaMcMlMlLm n ,,- .^^$ % //0@0@AB C"889 :	 r3   c                   dt        | j                         d}| j                  r$| j                  j                  j                         nd}|d| dz  }|d| j	                         j
                   dz  }|d| j                   dz  }|d| j                   dz  }|d| j                  | j                         dz  }|d	| j                   dz  }|S )
NzName: 
r+   z    Source: z    Create Function: z    Guard Types: z    Code List: z    Object Weakref: z    Guarded Class Weakref: )r   r   r   r   r   rL   r   r   r   r   r   )r/   outputr   s      r1   r2   zGuard.__str__h  s    $tyy/*"--1[[!!'')bL++)$*>*>*@*I*I)J"MM%d&6&6%7r::ODNN#3266()<)<T=M=M)N(OrRR/0J0J/K2NNr3   c           	     R   	 | j                  ||       S # t        $ r t        j                  dt	        |       j                                | j                  rNt        j                  ddj                  | j                  j                         dd        j                                 w xY w)NzError while creating guard:
%szCreated at:
%sr+   )
r   r>   log	exceptionrI   rstripr   errorjoinformat)r/   builders     r1   createzGuard.creates  s    	>>'400 	MM;SY=M=M=OPzz		+RWWTZZ5F5F5H5M-N-U-U-WX		s
    BB&c                6    | j                   j                         S rk   )r   rw   rV   s    r1   rw   zGuard.is_specialized_nn_module|  s    {{3355r3   c                6    | j                   j                         S rk   )r   rn   rV   s    r1   rn   zGuard.is_fsdp_module  s    {{))++r3   c                6    | j                   j                         S rk   )r   r   rV   s    r1   r   zGuard.is_local  s    {{##%%r3   c                    t        | j                  t        j                        r"| j                  j                  }|j
                  S | j                  }|j
                  S rk   )r   r   r   r   r   rL   )r/   r   s     r1   create_fn_namezGuard.create_fn_name  sJ    dnni&7&78++I !!! I!!!r3   c                   | j                   sg | _         | j                   j                  |       | j                  |d fvrt        d| j                   d|       || _        | j                  s|| _        n| j                  j                  |       | j                  |d fv xs) t        | j                        xr | j                         d u }|st        d| j                   d|       || _        y )Nz1Guarded class id must be identical, or None, got z vs zHGuarded object must be identical, None or ephemeral (dead weakref), got )r   appendr   r.   r   extendr   callable)r/   
guard_typeguarded_classr   r   is_valids         r1   set_export_infozGuard.set_export_info  s    !D
+%%mT-BB 112$}oG  &3"~~&DNNN!!), d 33 +(() +  "d* 	
  ''([M;  'r3   rH   r<   )rH   ztuple[bool, int, int, str, int])r   r   rH   r   )rH   z(Callable[[GuardBuilderBase, Guard], Any]rG   rH   rY   )r   objectrH   rI   )r   r   rH   r   r   )
r   rI   r   r   r   	list[str]r   r   rH   None)rL   rM   rN   rO   r   r   r   r   r   r   r   r   r   r   r   r   propertyr   r   staticmethodr   r   r2   r   rw   rn   r   r   r   r6   r3   r1   r   r      s   & 88 %)K!("&I&!%K%?C<C&*E#*04J-4E:!OT!

 2" , , 4 4 $ $0
	6,&"&'&' 9&' 	&'
 &' 
&'r3   r   Tc                      e Zd Zy)GuardEnvExprNr   r6   r3   r1   r   r     s    r3   r   c                  *    e Zd ZU ded<   ded<   ddZy)DuplicateInputsr   input_source_ainput_source_bc                f    | j                   | j                  k(  rt        d| j                          y )Nz9input_source_a and input_source_b must be different, got )r   r   r.   rV   s    r1   __post_init__zDuplicateInputs.__post_init__  s>    $"5"55 **+-  6r3   NrH   r   )rL   rM   rN   rO   r  r6   r3   r1   r   r     s    r3   r   c                  "    e Zd ZU ded<   ded<   y)StorageOverlapzlist[Source]overlapping_sourcesnon_overlapping_sourcesNr   r6   r3   r1   r  r    s    %%))r3   r  c                  0    e Zd Zedd       Zedd       Zy)Checkpointablec                     y rk   r6   rV   s    r1   copy_graphstatezCheckpointable.copy_graphstate  s    $'r3   c                     y rk   r6   r/   states     r1   restore_graphstatez!Checkpointable.restore_graphstate  s    47r3   N)rH   r   )r  r   rH   r   )rL   rM   rN   r   r  r  r6   r3   r1   r	  r	    s    ' '7 7r3   r	  c                  4    e Zd ZU dZded<   ddZd	dZd
dZy)GuardsCheckpointStatezW
    The GuardCheckpointState - it is the T of Checkpointable[T] for GuardsContext
    OrderedSet[Guard]dynamo_guardsc                    || _         y rk   )r  )r/   r  s     r1   __init__zGuardsCheckpointState.__init__  s
    *r3   c                n    | j                   j                  |j                         }t        |      dk(  ry|S )z
        Produces a delta against another GuardsCheckpointState.

        Returns None if no delta is found, otherwise, return an OrderedSet() of mismatched
        Guard type objects.
        r   N)r  
differencer   r/   r   rs      r1   diffzGuardsCheckpointState.diff
  s3     ))%*=*=>q6Q;r3   c                J    t        |t              sy| j                  |      d u S NF)r   r  r  r   s     r1   __eq__zGuardsCheckpointState.__eq__  s$    %!67yy4''r3   N)r  r  rH   r   )r   r  rH   OrderedSet[Guard] | Noner   r   rH   r   )rL   rM   rN   __doc__rO   r  r  r  r6   r3   r1   r  r     s     %$+
(r3   r  c                  4    e Zd ZU i Zded<   ddZddZd	dZy)
ModuleContextCheckpointStatedict[str, torch.nn.Module]
nn_modulesc                    || _         y rk   r$  )r/   r$  s     r1   r  z%ModuleContextCheckpointState.__init__  s	    $r3   c                    t        | j                  j                               j                  t        |j                  j                                     }t	        |      dk(  ry|S )z
        Produces a delta against another ModuleContextCheckpointState.

        Returns None if no delta is found, otherwise, return a set() of mismatched
        module key names.
        r   N)setr$  keysr  r   r  s      r1   r  z!ModuleContextCheckpointState.diff"  sM     $$&'223u7G7G7L7L7N3OPq6Q;r3   c                J    t        |t              sy| j                  |      d u S r  )r   r"  r  r   s     r1   r  z#ModuleContextCheckpointState.__eq__.  $    %!=>yy4''r3   N)r$  r#  rH   r   )r   r"  rH   set[str] | Noner  )rL   rM   rN   r$  rO   r  r  r  r6   r3   r1   r"  r"    s    -/J*/%
(r3   r"  c                  $    e Zd ZddZddZddZy)ModuleContextc                    i | _         y rk   r&  rV   s    r1   r  zModuleContext.__init__5  s	    *,r3   c                >    t        t        | j                              S rk   )r"  dictr$  rV   s    r1   r  zModuleContext.copy_graphstate8  s    +D,ABBr3   c                t    t        |t              st        dt        |             |j                  | _        y )Nz+expected ModuleContextCheckpointState, got )r   r"  r.   typer$  r  s     r1   r  z ModuleContext.restore_graphstate;  s7    %!=> =d5k]K   **r3   Nr  )rH   r"  )r  r"  rH   r   )rL   rM   rN   r  r  r  r6   r3   r1   r.  r.  4  s    -C+r3   r.  c                  4    e Zd ZU i Zded<   ddZddZd	dZy)
GlobalContextCheckpointStatedict[str, tuple[Callable, Any]]global_statec                    || _         y rk   r7  )r/   global_statess     r1   r  z%GlobalContextCheckpointState.__init__F  s
    )r3   c                    t        | j                  j                               j                  t        |j                  j                                     }t	        |      dk(  ry|S )z
        Produces a delta against another GlobalContextCheckpointState.

        Returns None if no delta is found, otherwise, return a set() of mismatched
        global key names.
        r   N)r(  r7  r)  r  r   r  s      r1   r  z!GlobalContextCheckpointState.diffI  sO     !!&&()44S9K9K9P9P9R5STq6Q;r3   c                J    t        |t              sy| j                  |      d u S r  )r   r5  r  r   s     r1   r  z#GlobalContextCheckpointState.__eq__U  r+  r3   N)r:  r6  rH   r   )r   r5  rH   r,  r  )rL   rM   rN   r7  rO   r  r  r  r6   r3   r1   r5  r5  C  s    46L16*
(r3   r5  c                  0    e Zd ZdZh dZddZddZd	dZy)
GlobalContextzz
    This keeps track of the global torch state during tracing of a function.
    For example, torch.is_grad_enabled.
    >   grad_enabledautocast_enabledautocast_cpu_dtypeautocast_gpu_dtypeautocast_cpu_enabledautocast_cache_enabledc                    i | _         y rk   r9  rV   s    r1   r  zGlobalContext.__init__j  s
    =?r3   c                ,    t        | j                        S rk   )r5  r7  rV   s    r1   r  zGlobalContext.copy_graphstatem  s    +D,=,=>>r3   c                   t        |t              st        dt        |             |j                  | _        t        | j                        t        | j                        k(  r0t        | j                  j                               | j                  k(  s<t        dt        | j                  j                                d| j                         | j                  j                         D ]  \  }} ||        y )Nz+expected GlobalContextCheckpointState, got z Global state mismatch: got keys z, expected )
r   r5  r.   r3  r7  r   _supported_global_statesr(  r)  values)r/   r  r   argss       r1   r  z GlobalContext.restore_graphstatep  s    %!=> =d5k]K  "..!!"c$*G*G&HHD%%**,-1N1NN 23t7H7H7M7M7O3P2Q R 99:<  ++224 	JD$J	r3   Nr  )rH   r5  )r  r5  rH   r   )rL   rM   rN   r   rH  r  r  r  r6   r3   r1   r>  r>  [  s    
 @?r3   r>  c                  |    e Zd ZdddZddZddZddZddZddZddZ	d	d
d	 	 	 	 	 	 	 ddZ
ddZddZddZy)	GuardsSetNc                    |t               | _        n|| _        t        t              | _        | j                  D ]  }| j                  |        y rk   )r   innerr   listsource_to_guardstrack_guard_by_source)r/   rN  guards      r1   r  zGuardsSet.__init__  sI    =,6LDJDJ CNdBSZZ 	.E&&u-	.r3   c                X    |j                   }| j                  |   j                  |       y rk   )r   rP  r   )r/   rR  r   s      r1   rQ  zGuardsSet.track_guard_by_source  s&    ))f%,,U3r3   c                ,    t        | j                        S rk   )iterrN  rV   s    r1   __iter__zGuardsSet.__iter__      DJJr3   c                ,    t        | j                        S rk   )r   rN  rV   s    r1   __len__zGuardsSet.__len__  s    4::r3   c                F    t        | j                  |j                  z
        S rk   )rL  rN  r   s     r1   __sub__zGuardsSet.__sub__  s    ekk122r3   c                ,    t        | j                        S rk   )r   rN  rV   s    r1   __bool__zGuardsSet.__bool__  rW  r3   c                J    t               | _        t        t              | _        y rk   )r   rN  r   rO  rP  rV   s    r1   clearzGuardsSet.clear  s    \
 +D 1r3   Tr   )collect_debug_stackskipc                  || j                   v ry |r*|j                  t        j                  d|z         |_        |j                  t
        j                         |_        | j                   j                  |       | j                  |       y NrZ   )ra  )	rN  r   r   extractr   TracingContextextract_stackaddrQ  )r/   rR  r`  ra  s       r1   rg  zGuardsSet.add  sq     DJJ{{"/77QXF#-;;=E

u""5)r3   c                F    |D ]  }|D ]  }| j                  |d         y rc  )rg  )r/   othersogs       r1   updatezGuardsSet.update  s0     	$A $#$	$r3   c                2    t        | j                  |         S )z4Return all guards with the given originating_source.)rO  rP  )r/   r   s     r1   get_guards_for_sourcezGuardsSet.get_guards_for_source  s    D))&122r3   c                    ddl m t        fd| j                  D              | _        t	        t
              | _        | j                  D ]  }| j                  |        y)z.Delete all guards that contains a given sourcerZ   )is_from_sourcec              3  J   K   | ]  } |j                         r|  y wrk   )r   ).0rk  rp  r   s     r1   	<genexpr>z6GuardsSet.remove_guards_with_source.<locals>.<genexpr>  s%       
8L8Lf)UA 
s   ##N)_dynamo.sourcerp  r   rN  r   rO  rP  rQ  )r/   r   rR  rp  s    ` @r1   remove_guards_with_sourcez#GuardsSet.remove_guards_with_source  sS    2  
zz 
 


 !,D 1ZZ 	.E&&u-	.r3   rk   )rN  r  rH   r   )rR  r   rH   r   )rH   zIterator[Guard]r   )r   rL  rH   rL  r   r  )rR  r   r`  r   ra  r<   rH   r   )ri  z
set[Guard]rH   r   )r   r   rH   zlist[Guard])r   r   rH   r   )rL   rM   rN   r  rQ  rV  rY  r[  r]  r_  rg  rl  rn  ru  r6   r3   r1   rL  rL    sa    	.4 
3 2
 <@Q**48*GJ*	*$
3.r3   rL  c                  J    e Zd ZddZej
                  dd       ZddZd	dZy)
GuardsContextc                >    t               | _        g | _        d| _        y r  )rL  r  aotautograd_guardsskip_installrV   s    r1   r  zGuardsContext.__init__  s    (168"'r3   c              #  b   K   | j                   }d| _         	 d  || _         y # || _         w xY ww)NT)rz  )r/   olds     r1   skip_guard_installz GuardsContext.skip_guard_install  s1      	$ #DDs   /# /	,/c                R    t        t        | j                  j                              S rk   )r  r   r  rN  rV   s    r1   r  zGuardsContext.copy_graphstate  s    $Z0B0B0H0H%IJJr3   c                    t        |t              st        dt        |             t	        |j
                        | _        y )Nz$expected GuardsCheckpointState, got )r   r  r.   r3  rL  r  r  s     r1   r  z GuardsContext.restore_graphstate  s7    %!67 #GU}!UVV&u':':;r3   Nr  rH   Generator[None, None, None])rH   r  )r  r  rH   r   )	rL   rM   rN   r  
contextlibr   r}  r  r  r6   r3   r1   rw  rw    s,    (
 $ $K<r3   rw  c                      e Zd Ze	 	 	 	 	 	 dd       Zedd       Zedd       Zedd       Zedd       Zedd       Z	e	 	 	 	 	 	 dd       Z
e	 	 	 	 dd       Ze	 	 	 	 	 	 	 	 dd	       Ze	 	 	 	 	 	 dd
       Zy)HopSubgraphCachec                     y rk   r6   r/   fn_code
identifiers      r1   add_dynamo_installed_submodulez/HopSubgraphCache.add_dynamo_installed_submodule       r3   c                     y rk   r6   r/   r  s     r1   get_dynamo_installed_submodulesz0HopSubgraphCache.get_dynamo_installed_submodules  s    ORr3   c                     y rk   r6   r/   r  keys      r1   add_autograd_key_entryz'HopSubgraphCache.add_autograd_key_entry  s    NQr3   c                     y rk   r6   r/   r  s     r1   get_autograd_key_entryz'HopSubgraphCache.get_autograd_key_entry  s    JMr3   c                     y rk   r6   r  s      r1   add_proxy_dispatch_entryz)HopSubgraphCache.add_proxy_dispatch_entry  s    PSr3   c                     y rk   r6   r  s     r1   get_proxy_dispatch_entryz)HopSubgraphCache.get_proxy_dispatch_entry  s    LOr3   c                     y rk   r6   r/   r  schemas      r1   add_functionalize_schema_entryz/HopSubgraphCache.add_functionalize_schema_entry   r  r3   c                     y rk   r6   r/   r  s     r1   get_functionalize_schema_entryz/HopSubgraphCache.get_functionalize_schema_entry  s     *-r3   c                     y rk   r6   )r/   r  tangent_metadatagmods       r1   add_lazy_bwd_entryz#HopSubgraphCache.add_lazy_bwd_entry
  s     r3   c                     y rk   r6   r/   r  r  s      r1   get_lazy_bwd_entryz#HopSubgraphCache.get_lazy_bwd_entry  s     :=r3   Nr  r   r  rI   rH   r   r  r   rH   r   r  rI   r  r   rH   r   r  rI   rH   zCallable | Noner  r   r  ztorch._C.FunctionSchemarH   r   r  r   rH   ztorch._C.FunctionSchema | Noner  rI   r  tuple[object]r  torch.fx.GraphModulerH   r<   r  rI   r  r  rH   z.tuple[torch.fx.GraphModule | None, int | None])rL   rM   rN   r   r  r  r  r  r  r  r  r  r  r  r6   r3   r1   r  r    s(   -0	  R RQ QM MS SO O#:	  --	'- -  ( #	
 
  ==1>=	7= =r3   r  c                  b    e Zd ZU ded<   ded<   ded<   ded<   d	ed
<   ded<   ded<   dZded<   y)InvokeSubgraphReuseEntryrI   	body_namer  	body_gmodz!NestedCompileRegionOptions | Noners   	list[Any]subgraph_input_mappingr   single_tensor_outputzIlist[tuple[torch.Size, tuple[int, ...], torch.dtype, torch.device, bool]]output_metadatazlist[Source | None]arg_sourcesr   r<   num_user_outputsN)rL   rM   rN   rO   r  r6   r3   r1   r  r    sD    N##--    %$ cr3   r  c                  `    e Zd ZU ded<   ded<   dZded<    ej                  e      Zd	ed
<   y)InvokeSubgraphReuseConditionzlist[tuple[Any, object]]input_checksz2list[tuple[Source, GuardCheckSpec, object, Guard]]r   Nzpytree.TreeSpec | Nonetreespec)default_factoryzOrderedSet[Source]traced_sources)	rL   rM   rN   rO   r  dataclassesfieldr   r  r6   r3   r1   r  r  2  s9     +*
 ?> (,H$+
 *;):)::)VN&Vr3   r  c                     e Zd ZddZ	 	 	 	 	 	 ddZddZddZddZddZddZ		 	 	 	 	 	 ddZ
	 	 	 	 dd	Z	 	 	 	 	 	 	 	 dd
Z	 	 	 	 	 	 ddZddZddZ	 d	 	 	 	 	 	 	 	 	 ddZ	 	 	 	 	 	 d dZ	 	 	 	 	 	 d!dZ	 d	 	 	 	 	 	 	 	 	 d"dZy)#InvokeSubgraphCachec                    i | _         i | _        i | _        t        t              | _        t        t              | _        i | _        t        t              | _	        t        t              | _
        y rk   )autograd_cacheproxy_dispatch_cachefunctionalize_schema_cacher   rO  dynamo_installed_submodulesr1  lazy_bwd_cacheeffects_cachesubgraph_reuse_cachesubgraph_reuse_key_cacherV   s    r1   r  zInvokeSubgraphCache.__init__N  sk    359;!QS'FQRVFW(  	
  	  	!  	%r3   c                @    | j                   |   j                  |       y rk   )r  r   r  s      r1   r  z2InvokeSubgraphCache.add_dynamo_installed_submodulee  s     	((188Dr3   c                :    | j                   j                  |g       S rk   )r  getr  s     r1   r  z3InvokeSubgraphCache.get_dynamo_installed_submodulesj  s    //33GR@@r3   c                "    || j                   |<   y rk   )r  r  s      r1   r  z*InvokeSubgraphCache.add_autograd_key_entrym  s    *-J'r3   c                :    | j                   j                  |d       S rk   )r  r  r  s     r1   r  z*InvokeSubgraphCache.get_autograd_key_entryp  s    ""&&z488r3   c                "    || j                   |<   y rk   )r  r  s      r1   r  z,InvokeSubgraphCache.add_proxy_dispatch_entrys  s    03!!*-r3   c                :    | j                   j                  |d       S rk   )r  r  r  s     r1   r  z,InvokeSubgraphCache.get_proxy_dispatch_entryv  s    ((,,Z>>r3   c                "    || j                   |<   y rk   )r  r  s      r1   r  z2InvokeSubgraphCache.add_functionalize_schema_entryy  s     06'',r3   c                :    | j                   j                  |d       S rk   )r  r  r  s     r1   r  z2InvokeSubgraphCache.get_functionalize_schema_entry~  s     ..223==r3   c                ^    t        | j                  |         }||f| j                  |   |<   |S rk   )r   r  )r/   r  r  r  	num_gmodss        r1   r  z&InvokeSubgraphCache.add_lazy_bwd_entry  s:     ++J78	=A9<MJ'(89r3   c                ^    || j                   vry| j                   |   j                  |d      S )N)NN)r  r  r  s      r1   r  z&InvokeSubgraphCache.get_lazy_bwd_entry  s4     T000"":.223C\RRr3   c           	         | j                   j                  |d      x}r||k7  rt        d| d| d| d      || j                   |<   y)z>Store the effect types for a given invoke_subgraph identifier.NzPDifferent number of effects were found for invoke_subgraph call with identifier z,. 
Previously we had the following effects: z.
But now we have: .)r  r  r.   )r/   r  effectsprev_effectss       r1   add_effectszInvokeSubgraphCache.add_effects  sn    --11*dCC<C,&$,,6< 8@@L~ N((/y3  *1:&r3   c                :    | j                   j                  |d      S )zARetrieve the effect types for a given invoke_subgraph identifier.N)r  r  r  s     r1   get_effectszInvokeSubgraphCache.get_effects  s    !!%%j$77r3   c                    | j                   |   }t        |      |k\  rt        d| d| d      |j                  ||f       y )N1invoke_subgraph: exceeded maximum reuse entries () for function code a?  . This most likely means a guard keeps failing on every invocation, preventing subgraph reuse. Set TORCH_LOGS='+hierarchical_compile' to identify which guard is failing. If reuse is genuinely not possible and you need more cache entries, increase the limit via the max_reuse_entries argument to nested_compile_region().)r  r   RuntimeErrorr   )r/   r  	conditionentrymax_reuse_entriesentriess         r1   add_reuse_entryz#InvokeSubgraphCache.add_reuse_entry  s`     ++G4w<,,%&&:7) DIJ	 	 		5)*r3   c                    | j                   j                  |g       }t        |      D ];  \  }\  }} |||      s|dkD  r!|j                  d|j	                  |             |c S  y rp   )r  r  	enumerateinsertpop)r/   r  	evaluatorr  ir  r  s          r1   find_reuse_entryz$InvokeSubgraphCache.find_reuse_entry  sk     ++//<%.w%7 	!A!	5E*q5NN1gkk!n5	 r3   c                X    | j                   j                  |i       j                  |      S rk   )r  r  )r/   r  hash_keys      r1   find_reuse_entry_by_keyz+InvokeSubgraphCache.find_reuse_entry_by_key  s(     ,,00"=AA(KKr3   c                t    | j                   |   }t        |      |k\  r||vrt        d| d| d      |||<   y )Nr  r  zc (hash-key path). Increase the limit via the max_reuse_entries argument to nested_compile_region().)r  r   r  )r/   r  r  r  r  	key_caches         r1   add_reuse_entry_by_keyz*InvokeSubgraphCache.add_reuse_entry_by_key  s^     11':	y>..893L%&&:7) D+,  $	(r3   Nr  r  r  r  r  r  r  r  r  )r  rI   r  r(  rH   r   )r  rI   rH   z
set | None)ra   )
r  r   r  r  r  r  r  r<   rH   r   )r  r   r  zHCallable[[InvokeSubgraphReuseCondition, InvokeSubgraphReuseEntry], bool]rH   InvokeSubgraphReuseEntry | None)r  r   r  r<   rH   r  )
r  r   r  r<   r  r  r  r<   rH   r   )rL   rM   rN   r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r6   r3   r1   r  r  M  s   .EE-0E	E
A.94?66#:6	6
>>	'>
		 (	 #		
 
	SS1>S	7S
18 "#++ 0+ (	+
 + 
+*
 
) LL+.L	(L "#$$ $ (	$
 $ 
$r3   r  c                      e Zd ZddZddZy)HopDispatchSetCachec                2    ddl m} |t               i| _        y )Nr   )invoke_subgraph)'torch._higher_order_ops.invoke_subgraphr  r  hop_cache_map)r/   r  s     r1   r  zHopDispatchSetCache.__init__  s    K-/B/DEr3   c                >    || j                   vry | j                   |   S rk   )r  )r/   ops     r1   	get_cachezHopDispatchSetCache.get_cache  s$    T'''!!"%%r3   Nr  )r  ztorch._ops.HigherOrderOperatorrH   zHopSubgraphCache | None)rL   rM   rN   r  r   r6   r3   r1   r  r    s    F&r3   r  c                  \    e Zd Zedd       Zedd       Zd	dZed
d       Zedd       Zy)CompileContextc                 X    t         j                  t        d      t         j                  S )Nzcompile_context is not set)_TLScompile_contextr.   r6   r3   r1   r  zCompileContext.get  s&    ' !=>>###r3   c                 $    t        t        dd       S Nr  getattrr  r6   r3   r1   try_getzCompileContext.try_get      t.55r3   c                    |'t        |t              st        dt        |             || _        d| _        g | _        y )Nz*compile_id must be None or CompileId, got r   )r   r%   r.   r3  r@   rS   shape_env_guards)r/   r@   s     r1   r  zCompileContext.__init__  sF    !*Z*K <T*=M<NO  -7+-r3   c                 H    t         j                         } | y | j                  S rk   )r  r
  r@   rV   s    r1   current_compile_idz!CompileContext.current_compile_id  s"    %%'<r3   c                     t         j                         } | y | j                  y t        | j                  | j                        S rk   )r  r
  r@   rR   rS   rV   s    r1   current_trace_idzCompileContext.current_trace_id   s:    %%'<??"t55r3   N)rH   r  )rH   CompileContext | None)r@   rK   rH   r   )rH   rK   )rH   zTraceId | None)	rL   rM   rN   r   r  r
  r  r  r  r6   r3   r1   r  r    sU    $ $
 6 6.   6 6r3   r  c                  0    e Zd ZU dZded<   ded<   ded<   y)	InlinedCodeCachez8Cache for code-object-derived data used during inlining.r  instructionszdict[Any, int]indexofr   code_optionsN)rL   rM   rN   r   rO   r6   r3   r1   r  r  *  s    Br3   r  c                  L   e Zd ZdZedd       Zedd       ZddZddZee	dd              Z
edd       ZddZeej                  dd	              Zeej                  	 	 	 	 dd
              Zeej                  dd              Ze	 d	 	 	 	 	 	 	 	 	 dd       Zedd       Zy)re  z
    Provides the currently installed TracingContext, or None.

    Note that it is a staticmethod, and invocations outside of `with tracing()` (see below), are valid but
    will return None.
    c                 $    t        t        dd       S )Ntracing_contextr  r6   r3   r1   r
  zTracingContext.try_get;  r  r3   c                 H    t         j                         x} r| S t        d      )Nz<TracingContext.get() must be called within an ongoing trace.)re  r
  r  )ctxs    r1   r  zTracingContext.get?  s+     ((**3*JJ
 	
r3   c                   t               | _        t               | _        t	               | _        t               | _        t               | _        t               | _	        || _
        g | _        d | _        d | _        d | _        d | _        d | _        d | _        d | _        d | _        d | _        d| _        t-               | _        d| _        t3               | _        g | _        d | _        y r  )rw  guards_contextr.  module_contextr>  global_contextr1  previously_inlined_functionspreviously_cleaned_instructionsinlined_code_cache	fake_modeframe_summary_stackloc_in_frameloc_in_frame_positionsfw_metadataddp_optimizer_ctxaot_graph_nameparams_flatparams_flat_unwrap_subclassesparams_unwrapped_to_flat_indexoutput_strides#force_unspec_int_unbacked_size_liker   tensor_to_contextfakify_first_callr  hop_dispatch_set_cachetraced_codecudagraph_annotation)r/   r$  s     r1   r  zTracingContext.__init__G  s    +o+o+o<@F)?Cv,?Cv09AC  :><@#7;=A04-1?C*@D+ DH 490!8!:
 "'&9&;#+-)-!r3   c                    i | j                   _        | j                  j                          | j                  j                          | j
                  j                          y rk   )r   r7  r!  r_  r"  r#  rV   s    r1   r_  zTracingContext.clear~  sH     ,.())//1,,224%%'r3   c               +  V  K   i }t         j                         }| D ]  }t        ||      ||<    | j                         D ]  \  }}t	        |||        	 d  |j                         D ]  \  }}t	        |||        y # |j                         D ]  \  }}t	        |||        w xY wwrk   )re  r  r	  r;   setattr)kwargspriorr  r  vals        r1   patchzTracingContext.patch  s        " 	+C c*E#J	+  	#HCCc"	#	'!KKM 'SS#&'EKKM 'SS#&'s   AB)A? &B)?'B&&B)c                     t         j                         } | t        j                         S | j                  }| j
                  || j                         gz   }t        j                  j                  |      S rk   )re  r
  	tracebackStackSummaryr%  r&  _populate_loc_in_frame_summary	from_list)r/   r   s     r1   rf  zTracingContext.extract_stack  sh    %%'<))++(((T@@BCCE%%//66r3   c                *   | j                   t        d      | j                   \  }}}i }t        j                  dk\  r>| j                  2| j                  j
                  |d<   | j                  j                  |d<   t        j                  |||fddi|S )Nzloc_in_frame must not be None)r\   rd   colno	end_colnolookup_lineF)	r&  r.   sysversion_infor'  
col_offsetend_col_offsetr=  FrameSummary)r/   filenamelineno
frame_namer8  s        r1   r?  z-TracingContext._populate_loc_in_frame_summary  s    $ !@AA'+'8'8$&*!#w&4+F+F+R"99DDF7O"&"="="L"LF;%%
 	

 
 	
r3   c               #  
  K   t         j                         } t        j                  j                  j                  | dg       5  t        j                  j                  j                  | dd       5  t        j                  j                  j                  | dd       5  	 d  	 d d d        d d d        d d d        y # t        $ r}t        |d      sd |_         d }~ww xY w# 1 sw Y   ?xY w# 1 sw Y   CxY w# 1 sw Y   y xY ww)Nr%  r&  r'  
real_stack)	re  r  unittestmockr;  r   r>   hasattrrN  )tcrE   s     r1   clear_framezTracingContext.clear_frame  s      !MM&&r+@"E	MM&&r>4@	 MM&&r+CTJ	
	 	 	 	  $ q,/#'AL)	 	 	 	 	 	sl   A D,C7.,C+CB:!C+)C71	D:	C	C	C	CC($C++C4	0C77D <Dc              #    K   t         j                         }| |j                  j                  |        |j                  }|j
                  }d |_        d |_        	 d  	 | |j                  j                          ||_        ||_        y # t        $ r'}t        |d      s|j                         |_	         d }~ww xY w# | |j                  j                          ||_        ||_        w xY ww)NrN  )re  r  r%  r   r&  r'  r>   rQ  rf  rN  r  )frame_summaryrR  r|  old_positionsrE   s        r1   current_framezTracingContext.current_frame  s      !$""))-8oo11$(!
	6 (&&**,!BO(5B%  	1l+!//1	
 (&&**,!BO(5B%s6   AC-B  +C-	B;"B66B;;B> >,C**C-c               #     K   t         j                         } | d  y | j                  }g | _        	 | j                   || _        y # || _        w xY wwrk   )re  r
  r.  )rR  old_output_stridess     r1   report_output_stridesz$TracingContext.report_output_strides  sY     
 ##%:J..	3### 2B 2Bs   /AA  A	AANc                N    t         j                         }| ||f|_        ||_        y rk   )re  r  r&  r'  )rJ  rK  rL  	positionsrR  s        r1   set_current_loczTracingContext.set_current_loc  s(     !#VZ8$-!r3   c                 H    t         j                         } | y | j                  S rk   )re  r
  r3  )rR  s    r1   get_traced_codezTracingContext.get_traced_code  s"    ##%:~~r3   )rH   TracingContext | None)rH   re  )r$  FakeTensorMode | NonerH   r   r  )r8  r   rH   r  )rH   ztraceback.StackSummary)rH   ztraceback.FrameSummaryr  )rU  ztraceback.FrameSummary | NonerH   r  )rH   z:Generator[list[tuple[int, ...] | None] | None, None, None]rk   )
rJ  rI   rK  r<   rL  rI   r\  zdis.Positions | NonerH   r   )rH   zlist[CodeType] | None)rL   rM   rN   r   r   r
  r  r  r_  r   r;  rf  r?  r  rS  rW  rZ  r]  r_  r6   r3   r1   re  re  3  s@    6 6 
 
5.n( '  ' 7 7
&   > 646	$6  60 3  3 
 +/	
.
.
. 
. (	
.
 

. 
.  r3   re  c              #     K   t        t        dd       }| t        _        	 |  |t        _        y # |t        _        w xY wwr  )r	  r  r  )contextold_contexts     r1   r  r    s9      $ 148K"D+*{s   A 0 A =A c              #    K   t        t        dd      }| t        _        	 |  	 | F| j                  :| j                  j                  $| j                  j                  j                          |t        _        y# t        $ r)}t	        |d      s| | j                         |_         d}~ww xY w# | F| j                  :| j                  j                  $| j                  j                  j                          |t        _        w xY ww)z
    This function installs the passed in tracing context as a dynamic scoped
    global variable.

    Calls to TracingContext.get() while not under a `with tracing()` context
    will return None.
    r  NrN  )
r	  r  r  r>   rQ  rf  rN  r$  	shape_envcleanup)rc  rd  rE   s      r1   tracingrh     s      $ 148K"D+ !!-!!++7''//1*  q,'G,?"002AL !!-!!++7''//1*s5   DA9 AD9	B+$B&&B++B. .ADDc                     y rk   r6   r?   r8  s     r1   dataclass_with_cached_hashrk  =  s    HKr3   c                     y rk   r6   rj  s     r1   rk  rk  A  s     $'r3   c                (    dfd}| |S  ||       S )Nc                |    t        j                  | fi }| j                  dfd}d }||_        ||_        |S )Nc                n    t        | d      st        j                  | d |              | j                  S )Nr   )rQ  r   __setattr__r   )r/   old_hashs    r1   r   z:dataclass_with_cached_hash.<locals>.wrap.<locals>.__hash__O  s-    4)""4(4.A::r3   c                r     t        j                         }t         fd|D              } j                  |fS )Nc              3  d   K   | ]'  }|j                   st        |j                         ) y wrk   )initr	  r   )rr  fr/   s     r1   rs  zOdataclass_with_cached_hash.<locals>.wrap.<locals>.__reduce__.<locals>.<genexpr>Y  s"      Q1!&&qvv!6 Qs   00)r  fieldstupler   )r/   rv  field_valuess   `  r1   
__reduce__z<dataclass_with_cached_hash.<locals>.wrap.<locals>.__reduce__T  s4     !''-F  Q QQLNNL11r3   r   )r  r   r   ry  )	cls_innernew_clsr   ry  rq  r8  s       @r1   wrapz(dataclass_with_cached_hash.<locals>.wrapK  sD    ''	<V<%%	
	2 $'r3   )rz  type[T]rH   r}  r6   )r?   r8  r|  s    ` r1   rk  rk  G  s    * {9r3   c                      e Zd ZddZddZddZddZej                  dd       Z	e
dd       Zej                  dd       Z	 	 	 	 	 	 	 	 ddZdd	Zdd
ZddZdddZy)r   c                     yr  r6   rV   s    r1   is_dict_keyzSource.is_dict_keyj      r3   c                     yr  r6   rV   s    r1   is_ephemeralzSource.is_ephemeralm  r  r3   c                    t         rk   NotImplementedErrorr/   codegens     r1   reconstructzSource.reconstructp  s    !!r3   c                    t         )z
        Reconstructs the source into a string of Python code. This method should
        be eventually implemented for all subclasses of Source to achieve full
        coverage of python wrapper code generation.
        r  r  s     r1   reconstruct_pycodezSource.reconstruct_pycodes  s
     "!r3   c                    t         rk   r  rV   s    r1   r   zSource.guard_source{  s    !!r3   c                    t         )a  
        A template for the name of the source. Used to prevent code duplication between
        `name` and `get_value`.

        For non-ChainedSources, `name` and `get_value` use the returned string directly.

        For ChainedSources, `name` and `get_value` expect the return to be a format string
        with `{0}` present - `name` and `get_value` will apply different values to this function's
        returned format string.
        r  rV   s    r1   _name_templatezSource._name_template  s
     "!r3   c                    | j                   S rk   )r  rV   s    r1   r   zSource.name  s    """r3   c                P    | |v r||    S t        | j                  ||      }||| <   |S rk   )evalr  )r/   globalslocalscacher   s        r1   	get_valuezSource.get_value  s7     5=;T(('6:dr3   c                ^    | j                   t        j                  u rt        t	        | |      S rk   )r   rY   r   r  r   )r/   fns     r1   
make_guardzSource.make_guard  s(     4 44%%T2r3   c                6    | j                   j                         S rk   )r   rw   rV   s    r1   rw   zSource.is_specialized_nn_module  s      99;;r3   c                <    | j                   t        j                  k7  S )z+True if you can guard on attributes of this)r   rY   r   rV   s    r1   subguards_allowedzSource.subguards_allowed  s      K$?$???r3   Nc                    | ||       S | S rk   r6   )r/   transform_fns     r1   clonezSource.clone  s     #%%r3   r   )r  r   rH   r   )r  r   rH   rI   r   rG   r  dict[str, Any]r  r  r  zdict[Source, Any]rH   r   )r  zCallable[..., Any]rH   r   rk   r  z!Callable[[Source], Source] | NonerH   r   )rL   rM   rN   r  r  r  r  r   cached_propertyr   r   r  r   r  r  rw   r  r  r6   r3   r1   r   r   h  s    "" " " " " # #

 
 !	

 


<@r3   r   c                      e Zd ZU ded<   ddZddZej                  dd       ZddZ	ej                  dd       Z
	 	 	 	 	 	 	 	 ddZddd
Zy	)ChainedSourcer   basec                6    | j                   j                         S rk   )r  r  rV   s    r1   r  zChainedSource.is_dict_key  s    yy$$&&r3   c                6    | j                   j                         S rk   )r  r  rV   s    r1   r  zChainedSource.is_ephemeral  s    yy%%''r3   c                .    | j                   j                  S rk   )r  r   rV   s    r1   r   zChainedSource.guard_source  s    yy%%%r3   c                d    | }t        |t              r|j                  }t        |t              r|S rk   )r   r  r  )r/   currents     r1   get_basezChainedSource.get_base  s+    -0llG -0r3   c                `    | j                   j                  | j                  j                        S rk   )r  r   r  r   rV   s    r1   r   zChainedSource.name  s!    ""))$))..99r3   c                    | |v r||    S d}d}||v rd| }|dz  }||v r| j                   j                  |||      ||<   t        | j                  j	                  |      ||      }||= ||| <   |S )Ntmpr   rZ   )r  r  r  r  r   )r/   r  r  r  tmpvarcounterr   s          r1   r  zChainedSource.get_value  s     5=;7)_FqLG  ,,WfeDvT((//7&I6Ndr3   Nc                ^   d| j                   j                  |      i}t        j                  |       D ]W  }|j                  dk(  rt        | |j                        }t        |t              s:|j                  |      ||j                  <   Y t        j                  | fi |}| ||      }|S )Nr  )	r  r  r  rv  r   r	  r   r   replace)r/   r  cloned_fieldsru  r:  results         r1   r  zChainedSource.clone  s    )/1N(O##D) 	@Avv$'C#v&(+		,(?aff%	@ $$T;];#!&)Fr3   r   r   )rH   r   rG   r  rk   r  )rL   rM   rN   rO   r  r  r   r  r   r  r   r  r  r6   r3   r1   r  r    sz    
L'( & & : :  !	
 
&r3   r  c                X   ddl m}m}m} t        j                         x}r|j                  }||S g }ddlm} t        t         |                   D ]&  \  }}	t        |	|      s|j                  |	d|f       ( t        j                  |       }
t        |
      D ]  \  }}t        ||      r|j                  |j                  d|f       t        |      s<g } |||       |D cg c]  }t        ||      s| }}|j!                  t        |      D cg c]  \  }}|j                  d| |f c}}        |r`|d   \  }}}|d	d D ]M  \  }	}}||	ust#        d
| d| d| d|	 d| d| d| d| d|j$                   d| d| d|	j$                          |S yc c}w c c}}w )a  
    Attempts to "detect" what the current fake mode is.  If there is one ambiently
    available from TracingContext, we preferentially use that.  Otherwise, we
    heuristically detect the fake mode via the following sources, in order of
    priority:

        - Currently active fake mode on stack
        - Fake mode associated with passed in tensors (inputs does not
          have to be flattened)
    r   )
FakeTensorr    get_plain_tensorsN _get_current_dispatch_mode_stackzactive fake modezfake tensor input)outzsubclass input rZ   zfake mode (z) from r   z doesn't match mode (z

fake mode from z allocated at:
z
fake mode from )torch._subclasses.fake_tensorr  r    r  re  r
  r$  torch.utils._python_dispatchr  r  reversedr   r   pytreetree_leavesr   r   r.   r   )inputsr  r    r  rc  r$  
fake_modesr  r  mflat_inputs
flat_inputr  xfake_tensorsixtensordesc1i1desc2i2s                        r1   detect_fake_moder    s     !((**w*%%	 JM(#C#EFG :1a(q"4a89: $$V,K";/ :j*-z335H!LM(4;=Cjc2.*Q
";.L .  '0&="F %%'<bA  )!}	5"&qrN 	LAub!$!)GE7!B4?TUVTWW^_d^eefgifj k&&+WAbT1A)//AR S&&+WAbT1A!''L 	 -.s   'F!9F!F&c                 ~    ddl m}  ddlm} t	        t         |                   D ]  \  }}t        ||       s|c S  y)z~
    Inspects the dispatch mode stack for an active fake mode and returns it.
    Returns None if no fake mode is active.
    r   r   r  N)r  r    r  r  r  r  r   )r    r  rU   r  s       r1   active_fake_moder  -  s?    
 =M(#C#EFG 1a(H r3   )rc  r  rH   z,Generator[CompileContext | None, None, None])rc  r`  rH   z,Generator[TracingContext | None, None, None])r?   r}  r8  r   rH   r}  rk   )r?   r   r8  r   rH   zCallable[[type[T]], type[T]])r?   ztype[T] | Noner8  r   rH   z&type[T] | Callable[[type[T]], type[T]])r  r   rH   ra  )rH   ra  )k
__future__r   r  r  enumr   loggingrerE  	threadingr=  unittest.mockrO  r   abcr   collectionsr   r   r   typingr   r	   r
   r   r   r   typing_extensionsr   r   torch.utilsr   r  torch.utils._ordered_setr   r  r   torch.utils._tracebackr   r   torch.utils.weakr   	getLoggerrL   r   discollections.abcr   r   r   typesr   sympy"torch._dynamo.backends.distributedr   torch._dynamo.codegenr   torch._dynamo.guardsr   torch._dynamo.output_graphr   &torch._functorch._aot_autograd.schemasr   r  r   r  r    compiler7   r8   r%   rR   EnumrY   r   r   r   r   r   r   r   r  r	  r  r"  r.  r5  r>  rL  rw  r  r  r  r  r  localr  r  r  re  r  rh  rk  r   r  r  r  r6   r3   r1   <module>r     s   "      	 
      # % ! M M 1  ) / F B 4 g! ==F/36JR<  RZZ PQ "

X " $D17J 7J 27Jt7j 7<
$)) <
~	 	 d#  $$  T"u' u' #u'p CL d#	 	 $	 d#	l 	 $	 d#*\ * $*

8WQZ 8( (8( (0+N#?@ +( (0$N#?@ $RC. C.L<N#89 <2*= *=Z   2 W W W4R$* R$j
& 
& y&#6 #6L   ^ ^B +"+1+ + +"+1+ +8 
 K 
 K 
'	' #'!' 
'
 	*-+ @ 4(E E )ER 4(6F 6 )6r>Br3   