+
    &jV                      a  0 t $ ^ RIHt ^ RIt^ RIHt ^ RIHtHtHt ^ RI	H
t
HtHt ^ RIt^ RIHt ]'       d   ^ RIHt ^RIHt . R!Ot]! R4      t]
! R4      t]! ]P.                  R4      '       g^   ]! R4      ]P.                  P0                  R&   ]! R4      ]P.                  P0                  R&   ]! R4      ]P.                  P0                  R&   ^ RIHtHtHt R R ltR R lt ! R R	]4      t ! R R
4      t ]PB                  PD                  ]R]#3,          ,          t$R]%R&   ]R"R R ll4       t&]R"R R ll4       t&R"R R  llt&R# )#    )annotationsN)Callable)overloadTYPE_CHECKING	TypeAlias)	ParamSpecSelfTypeVar)Tensor)_POOL_HANDLE)_dummy_typeXPUGraphgraph_R_P_XpuStreamBase	_XPUGraph_xpu_graph_pool_handle_xpu_isCurrentStreamCapturing)r   r   r   c                   V ^8  d   QhRR/# )   returnbool )formats   "h/Users/jameslopez/projects/CWCArchive/cwc-podcast/.venv/lib/python3.14/site-packages/torch/xpu/graphs.py__annotate__r   (   s     + +T +    c                     \        4       # )zReturn True if XPU graph capture is underway on the current XPU stream, False otherwise.

If a XPU context does not exist on the current device, returns False without initializing the context.
)r   r   r   r   is_current_stream_capturingr    (   s    
 )**r   c                   V ^8  d   QhRR/# r   r   r   r   )r   s   "r   r   r   0   s     < << <r   c                 P    \         P                  P                  \        4       4      # )zBReturn an opaque token representing the id of a graph memory pool.)torchxpur   r   r   r   r   graph_pool_handler&   0   s    99!!"8":;;r   c                     a  ] tR t^5tRtRR V 3R llltRR V 3R llltR V 3R lltR V 3R	 lltR
 V 3R llt	R V 3R llt
R V 3R lltR V 3R lltR V 3R lltR V 3R lltR V 3R lltRtV ;t# )r   a  Wrapper around a XPU graph.

Arguments:
    keep_graph (bool, optional): If ``keep_graph=False``, the
        executable command graph will be instantiated on GPU at the end of
        ``capture_end`` and the underlying modifiable command graph will be
        destroyed. Note that the executable command graph will not be
        instantiated at the end of ``capture_end`` in this
        case. Instead, it will be instantiated via an explicit called
        to ``instantiate`` or automatically on the first call to
        ``replay`` if ``instantiate`` was not already called. Calling
        ``instantiate`` manually before ``replay`` is recommended to
        prevent increased latency on the first call to ``replay``.

c                    V ^8  d   QhRRRR/# )r   
keep_graphr   r   r	   r   )r   s   "r   r   XPUGraph.__annotate__F   s     0 0 0$ 0r   c                	"   < \         SV `  W4      # N)super__new__)clsr)   	__class__s   &&r   r.   XPUGraph.__new__F   s    ws//r   c                    V ^8  d   QhRRRR/# )r   pool_POOL_HANDLE | Noner   Noner   )r   s   "r   r   r*   I   s     ) )"5 ) )r   c                (   < \         SV `  VR7       R# )a  Begin capturing XPU work on the current xpu stream.

Typically, you shouldn't call ``capture_begin`` yourself.
Use :class:`~torch.xpu.graph`, which call ``capture_begin`` internally.

Arguments:
    pool (optional): Token (returned by :func:`~torch.xpu.graph_pool_handle` or
        :meth:`other_Graph_instance.pool()<torch.xpu.XPUGraph.pool>`) that hints this graph may share memory
        with the indicated pool.
r3   N)r-   capture_begin)selfr3   r0   s   &&r   r8   XPUGraph.capture_beginI   s     	4(r   c                   V ^8  d   QhRR/# r   r   r5   r   )r   s   "r   r   r*   V   s      T r   c                $   < \         SV `  4        R# )zEnd XPU graph capture on the current stream.

After ``capture_end``, ``replay`` may be called on this instance.

Typically, you shouldn't call ``capture_end`` yourself.
Use :class:`~torch.xpu.graph`, which call ``capture_end`` internally.
N)r-   capture_endr9   r0   s   &r   r>   XPUGraph.capture_endV   s     	r   c                   V ^8  d   QhRR/# r<   r   )r   s   "r   r   r*   `   s      T r   c                $   < \         SV `  4        R# )a  Instantiate the XPU graph. Will be called by
``capture_end`` if ``keep_graph=False``, or by ``replay`` if
``keep_graph=True`` and ``instantiate`` has not already been
explicitly called. Does not destroy the xpu modify command graph returned
by ``raw_xpu_graph``.
N)r-   instantiater?   s   &r   rC   XPUGraph.instantiate`   s     	r   c                   V ^8  d   QhRR/# r<   r   )r   s   "r   r   r*   i   s       r   c                $   < \         SV `  4        R# )z+Replay the XPU work captured by this graph.N)r-   replayr?   s   &r   rG   XPUGraph.replayi   s    r   c                   V ^8  d   QhRR/# r<   r   )r   s   "r   r   r*   m   s      t r   c                $   < \         SV `  4        R# )z1Delete the graph currently held by this instance.N)r-   resetr?   s   &r   rK   XPUGraph.resetm   s    r   c                   V ^8  d   QhRR/# r"   r   )r   s   "r   r   r*   q   s      l r   c                    < \         SV `  4       # )zReturn an opaque token representing the id of this graph's memory pool.

This id can optionally be passed to another graph's ``capture_begin``,
which hints the other graph may share the same memory pool.
)r-   r3   r?   s   &r   r3   XPUGraph.poolq   s     w|~r   c                   V ^8  d   QhRR/# r<   r   )r   s   "r   r   r*   y   s     + +4 +r   c                    < \         SV `  4       # )z.Enable debugging mode for XPUGraph.debug_dump.)r-   enable_debug_moder?   s   &r   rR   XPUGraph.enable_debug_modey   s    w(**r   c                    V ^8  d   QhRRRR/# )r   
debug_pathstrr   r5   r   )r   s   "r   r   r*   }   s     . .S .T .r   c                "   < \         SV `  V4      # )z
Arguments:
    debug_path (required): Path to dump the graph to.

Calls a debugging function to dump the graph if the debugging is
enabled via XPUGraph.enable_debug_mode()
)r-   
debug_dump)r9   rU   r0   s   &&r   rX   XPUGraph.debug_dump}   s     w!*--r   c                   V ^8  d   QhRR/# r   r   intr   )r   s   "r   r   r*      s     ' 's 'r   c                    < \         SV `  4       # )zuReturns the underlying xpuGraph_t. ``keep_graph`` must be True.

XPU doesn't provide APIs to manipulate this object.
)r-   raw_xpu_graphr?   s   &r   r^   XPUGraph.raw_xpu_graph   s    
 w$&&r   c                   V ^8  d   QhRR/# r[   r   )r   s   "r   r   r*      s     , ,C ,r   c                    < \         SV `  4       # )a  Returns the underlying xpuGraphExec_t. ``instantiate`` must have been called if ``keep_graph`` is True, or ``capture_end`` must have been called if ``keep_graph`` is False. If you call ``instantiate()`` after ``raw_xpu_graph_exec()``, the previously returned xpuGraphExec_t will be destroyed. It is your responsibility not to use this object after destruction.

XPU doesn't provide APIs to manipulate this object.
)r-   raw_xpu_graph_execr?   s   &r   rb   XPUGraph.raw_xpu_graph_exec   s    
 w)++r   r   )Fr,   )__name__
__module____qualname____firstlineno____doc__r.   r8   r>   rC   rG   rK   r3   rR   rX   r^   rb   __static_attributes____classcell__)r0   s   @r   r   r   5   sv     0 0) )     + +. .' ', ,r   c                  R    ] tR t^t$ RtRtR]R&   RR R lltR R ltR	 R
 lt	Rt
R# )r   a^  Context-manager that captures XPU work into a :class:`torch.xpu.XPUGraph` object for later replay.

Arguments:
    xpu_graph (torch.xpu.XPUGraph): Graph object used for capture.
    pool (optional): Opaque token (returned by a call to :func:`~torch.xpu.graph_pool_handle()` or
        :meth:`other_Graph_instance.pool()<torch.xpu.XPUGraph.pool>`) hinting this graph's capture
        may share memory from the specified pool.
    stream (torch.xpu.Stream, optional): If supplied, will be set as the current stream in the context.
        If not supplied, ``graph`` sets its own internal side stream as the current stream in the context.

.. note::
    For effective memory sharing, if you pass a ``pool`` used by a previous capture and the previous capture
    used an explicit ``stream`` argument, you should pass the same ``stream`` argument to this capture.

Ntorch.xpu.Stream | Nonedefault_capture_streamc               $    V ^8  d   QhRRRRRR/# )r   	xpu_graphr   r3   r4   streamrl   r   )r   s   "r   r   graph.__annotate__   s(     # ## "# (	#r   c                	N   V P                   P                  f.   \        P                  P	                  4       V P                   n        Vf   RMV3V n        Ve   TMV P                   P                  V n        V P                  f   \        R4      hV P                  V n        Wn	        R # )Nzcapture_stream must not be Noner   )
r0   rm   r$   r%   Streamr3   capture_streamAssertionError
stream_ctxro   )r9   ro   r3   rp   s   &&&&r   __init__graph.__init__   s     >>00849II4D4D4FDNN1;?<RdW	(Fdnn.S.S 	 & !BCC--"r   c                   V ^8  d   QhRR/# r<   r   )r   s   "r   r   rq      s     1 14 1r   c                	    \         P                  P                  4        \         P                  P                  4        V P                  P                  4        V P                  P                  ! V P                  !   R # r,   )	r$   r%   synchronizeempty_cacherv   	__enter__ro   r8   r3   )r9   s   &r   r}   graph.__enter__   sH    				!!#$$dii0r   c                    V ^8  d   QhRRRR/# )r   argsobjectr   r5   r   )r   s   "r   r   rq      s     ( (f ( (r   c                	n    V P                   P                  4        V P                  P                  ! V!   R # r,   )ro   r>   rv   __exit__)r9   r   s   &*r   r   graph.__exit__   s$    ""$  $'r   )rt   r3   rv   ro   )NN)rd   re   rf   rg   rh   rm   __annotations__rw   r}   r   ri   r   r   r   r   r      s)      7;3:#*1( (r   .r   _ModuleOrCallablec               0    V ^8  d   QhRRRRRRRRR	R
RR/# )r   	callablesr   sample_argstuple[Tensor, ...]num_warmup_itersr\   allow_unused_inputr   r3   r4   r   r   )r   s   "r   r   r      sD       #  	
  r   c                    R # r,   r   r   r   r   r   r3   s   &&&&&r   make_graphed_callablesr      s     r   c               0    V ^8  d   QhRRRRRRRRR	R
RR/# )r   r   ztuple[_ModuleOrCallable, ...]r   ztuple[tuple[Tensor, ...], ...]r   r\   r   r   r3   r4   r   r   )r   s   "r   r   r      sD     ( (,(/( ( 	(
 ( #(r   c                    R # r,   r   r   s   &&&&&r   r   r      s     %(r   c               0    V ^8  d   QhRRRRRRRRR	R
RR/# )r   r   z1_ModuleOrCallable | tuple[_ModuleOrCallable, ...]r   z3tuple[Tensor, ...] | tuple[tuple[Tensor, ...], ...]r   r\   r   r   r3   r4   r   r   )r   s   "r   r   r      sL     n n@nDn n 	n
 n 7nr   c                   \         P                  ! 4       '       d'   \         P                  ! 4       '       d   \        R4      hRp\	        V \
        4      '       g0   RpV 3p \        P                  ! \
        \        R3,          V4      3pM5\        P                  ! \
        \
        \        R3,          R3,          V4      p. p\        W4       EFq  w  r\	        V\         P                  P                  4      '       d   \        VP                  4      ^ 8X  d5   \        VP                  4      ^ 8X  d   \        VP                  4      ^ 8X  g   \        R4      h\         ;QJ d*    R VP#                  4        4       F  '       d   K   RM	  RM! R VP#                  4        4       4      '       g   \        R4      h\         P$                  P&                  P(                  ! V	!  p
VP+                  \        V
4      4       \         ;QJ d    R V
 4       F  '       d   K   RM	  RM! R V
 4       4      '       d   EKi  \-        R4      h	  V U	u. uF  p	\        V	4      NK  	  pp	V  Uu. uFH  p\	        V\         P                  P                  4      '       d   \        VP/                  4       4      MRNKJ  	  pp\1        \        V 4      4       Uu. uF  pW},          W,          ,           NK  	  pp\1        \        V 4      4       Uu. uF!  p\         P2                  P5                  4       NK#  	  pp\1        \        V 4      4       Uu. uF!  p\         P2                  P5                  4       NK#  	  ppVf   \7        4       MTp\         P2                  P9                  4        \         P2                  P;                  \         P2                  P=                  4       4      ;_uu_ 4        \        WV4       EF  w  pp	pRRRppp\1        V4       F  p\         P$                  P&                  P?                  V! V	!  4      p\
        ;QJ d    . R V 4       F  NK  	  5M! R V 4       4      p\        V4      ^ 8  g   Kn  \         P@                  PC                  T\
        ;QJ d    . R	 V 4       F  NK  	  5M! R	 V 4       4      \
        ;QJ d    . R
 V 4       F  NK  	  5M! R
 V 4       4      RVR7      pK  	  VVV3 F  p?K  	  EK  	  RRR4       \         P2                  P9                  4        . p. p\        WV4       F  w  pp	p\         P2                  PE                  VVR7      ;_uu_ 4        V! V	!  pRRR4       \         P$                  P&                  PG                  X4      w  ppVP+                  \        V4      4       VP+                  V4       K  	  . p. p \        \I        V4      \I        V4      \I        V4      4       EF  w  pp!p"\
        ;QJ d    . R V! 4       F  NK  	  5M! R V! 4       4      p#\
        ;QJ d    . R V! 4       F  NK  	  5M! R V! 4       4      pRp\        V4      ^ 8  d   \         P2                  PE                  V"VR7      ;_uu_ 4        \         P@                  PC                  T\
        ;QJ d    . R V 4       F  NK  	  5M! R V 4       4      \
        ;QJ d    . R V# 4       F  NK  	  5M! R V# 4       4      RVR7      pRRR4       . p$^ p%V FM  p&V&PJ                  '       d(   Ve$   V$P+                  VV%,          4       V%^,          p%K<  V$P+                  R4       KO  	  \        V$4      p$VP+                  V#4       V P+                  V$4       EK  	  VPM                  4        V PM                  4        R R lp'. p(\O        V 4       F  w  ppV'! VV,          VV,          W,          W,          VV,          W,          VV,          VV,          V V,          4	      p)\	        V\         P                  P                  4      '       d>   R R lp*V*! VVPP                  V)VPR                  4      Vn)        V(P+                  V4       K  V(P+                  V)4       K  	  V'       d
   V(^ ,          # \        V(4      # u up	i u upi u upi u upi u upi   + '       g   i     EL; i  + '       g   i     EL; i  + '       g   i     EL; i)a  Accept callables (functions or :class:`nn.Module<torch.nn.Module>`\ s) and returns graphed versions.

Each graphed callable's forward pass runs its source callable's
forward XPU work as a XPU graph inside a single autograd node.

The graphed callable's forward pass also appends
a backward node to the autograd graph. During backward, this node runs the
callable's backward work as a XPU graph.

Therefore, each graphed callable should be a drop-in replacement for its source callable
in an autograd-enabled training loop.

See :ref:`Partial-network capture<partial-network-capture>` for detailed use and constraints.

If you pass a tuple of several callables, their captures will use the same memory pool.

Arguments:
    callables (torch.nn.Module or Python function, or tuple of these): Callable or callables to graph.
        If you pass a tuple of callables, their order in the tuple must be the same order they'll run
        in the live workload.
    sample_args (tuple of Tensors, or tuple of tuples of Tensors): Samples args for each callable.
        If a single callable was passed, ``sample_args`` must be a single tuple of argument Tensors.
        If a tuple of callables was passed, ``sample_args`` must be tuple of tuples of argument Tensors.
    num_warmup_iters (int): The number of warmup iterations. Currently, ``DataDistributedParallel`` needs
        11 iterations for warm up. Default: ``3``.
    allow_unused_input (bool): If False, specifying inputs that were not used when computing outputs
        (and therefore their grad is always zero) is an error. Defaults to False.
    pool (optional): Token (returned by :func:`~torch.xpu.graph_pool_handle` or
        :meth:`other_Graph_instance.pool()<torch.xpu.XPUGraph.pool>`) that hints this graph may share memory
        with the indicated pool.
.. note::
    The ``requires_grad`` state of each Tensor in ``sample_args`` must match the state
    that's expected for the corresponding real input in the training loop.

.. warning::
    This API is in beta and may change in future releases.

.. warning::
    ``sample_args`` for each callable must contain only Tensors. Other types are not allowed.

.. warning::
    Returned callables do not support higher order differentiation (e.g., double backward).

.. warning::
    In any :class:`~torch.nn.Module` passed to :func:`~make_graphed_callables`, only parameters
    may be trainable. Buffers must have ``requires_grad=False``.

.. warning::
    After you pass a :class:`torch.nn.Module` through :func:`~make_graphed_callables`,
    you may not add or remove any of that Module's parameters or buffers.

.. warning::
    :class:`torch.nn.Module`\s passed to :func:`~torch.xpu.make_graphed_callables` must not have module hooks
    registered on them at the time they are passed. However, registering hooks on modules *after* passing them
    through :func:`~torch.xpu.make_graphed_callables` is allowed.

.. warning::
    When running a graphed callable, you must pass its arguments in the same order and format
    they appeared in that callable's ``sample_args``.

.. warning::
    The automatic mixed precision is supported in :func:`~torch.xpu.make_graphed_callables` only with disabled
    caching. The context manager `torch.amp.autocast()` must have `cache_enabled=False`.
z_make_graphed_callables does not support the autocast caching. Please set `cache_enabled=False`.FT.c              3  <   "   T F  qP                   R J x  K  	  R# 5i)FNrequires_grad.0bs   & r   	<genexpr>)make_graphed_callables.<locals>.<genexpr>F  s     EA%/s   c              3  V   "   T F  p\        V\        P                  4      x  K!  	  R # 5ir,   )
isinstancer$   r   )r   args   & r   r   r   N  s     HKS:c5<<00Ks   ')Nc              3  L   "   T F  qP                   '       g   K  Vx  K  	  R # 5ir,   r   r   os   & r   r   r   n  s     $K1??QQ   $
$c              3  L   "   T F  qP                   '       g   K  Vx  K  	  R # 5ir,   r   r   is   & r   r   r   r  s      %';!AA';r   c              3  t   "   T F.  qP                   '       g   K  \        P                  ! V4      x  K0  	  R # 5ir,   r   r$   
empty_liker   s   & r   r   r   u  s(      +9@AOO/E,,Q//s   88)outputsinputsgrad_outputsonly_inputsallow_unusedr7   c              3  t   "   T F.  qP                   '       d   \        P                  ! V4      MR x  K0  	  R # 5ir,   r   r   s   & r   r   r     s)      $
FT???EQ<ns   68c              3  L   "   T F  qP                   '       g   K  Vx  K  	  R # 5ir,   r   r   s   & r   r   r     s     J1//QQr   c              3  L   "   T F  qP                   '       g   K  Vx  K  	  R # 5ir,   r   r   s   & r   r   r     s      T,@qOO,@r   c              3  0   "   T F  qf   K  Vx  K  	  R # 5ir,   r   r   s   & r   r   r     s     &W2EQqq2Es   
c               @    V ^8  d   QhRRRRRRRRRR	R
RRRRRRRRR/
# )r   	fwd_graphr   	bwd_graphmodule_paramsztuple[torch.nn.Parameter, ...]len_user_argsr\   output_unflatten_specztorch.utils._pytree.TreeSpecstatic_input_surfacer   static_outputsstatic_grad_outputsztuple[Tensor | None, ...]static_grad_inputsr   zCallable[..., object]r   )r   s   "r   r   ,make_graphed_callables.<locals>.__annotate__  sl     2 222 62 	2
  <2 12 +2 72 /2 
2r   c	           	        a aaaaaaaaa
  ! VV VVVVV3R  lR\         P                  P                  4      o
R V
VV3R llp	V	# )c                     < ] tR tRt]R VVVV3R ll4       t]]P                  P                  P                  R V VV3R ll4       4       t
RtR# )Omake_graphed_callables.<locals>.make_graphed_autograd_function.<locals>.Graphedi  c               $    V ^8  d   QhRRRRRR/# )r   ctxr   r   r   r   r   r   )r   s   "r   r   \make_graphed_callables.<locals>.make_graphed_autograd_function.<locals>.Graphed.__annotate__  s'     A AV Af A9K Ar   c                	  < \        S4       FS  pSV,          P                  4       W,          P                  4       8w  g   K5  SV,          P                  W,          4       KU  	  SP                  4        \	        S\
        4      '       g   \        R 4      h\
        ;QJ d    . R S 4       F  NK  	  5# ! R S 4       4      # )zstatic_outputs must be a tuplec              3  @   "   T F  qP                  4       x  K  	  R # 5ir,   detachr   s   & r   r   jmake_graphed_callables.<locals>.make_graphed_autograd_function.<locals>.Graphed.forward.<locals>.<genexpr>  s     @AXXZZs   )rangedata_ptrcopy_rG   r   tupleRuntimeError)r   r   r   r   r   r   r   s   &* r   forwardWmake_graphed_callables.<locals>.make_graphed_autograd_function.<locals>.Graphed.forward  s     }-A+A.779VY=O=O=QQ,Q/55fi@ .   "!.%88&'GHHu@@u@u@@@@r   c               $    V ^8  d   QhRRRRRR/# )r   r   r   gradsr   r   r   r   )r   s   "r   r   r     s"      f f 9K r   c                	  < \        V4      \        S4      8w  d$   \        R \        S4       R\        V4       24      h\        SV4       FA  w  r#Vf   K  VP                  4       VP                  4       8w  g   K0  VP	                  V4       KC  	  SP                  4        \        S\        4      '       g   \        R4      h\        ;QJ d    . R S 4       F  NK  	  5# ! R S 4       4      # )z	Expected z gradients but got z"static_grad_inputs must be a tuplec              3  L   "   T F  qe   VP                  4       MTx  K  	  R # 5ir,   r   r   s   & r   r   kmake_graphed_callables.<locals>.make_graphed_autograd_function.<locals>.Graphed.backward.<locals>.<genexpr>  s!      @R1-AHHJQ6@Rs   "$)lenr   zipr   r   rG   r   r   )r   r   ggradr   r   r   s   &*  r   backwardXmake_graphed_callables.<locals>.make_graphed_autograd_function.<locals>.Graphed.backward  s     u:%8!99&#C(;$<#==PQTUZQ[P\]   ##6>GA}::<4==?:GGDM  ?   "!"4e<<&'KLLu @Ru u @R  r   r   N)rd   re   rf   rg   staticmethodr   r$   autogradfunctiononce_differentiabler   ri   )r   r   r   r   r   r   r   s   r   Graphedr     sN    A A A ^^$$88  9 r   r   c                    V ^8  d   QhRRRR/# )r   	user_argsr   r   r   )r   s   "r   r   Tmake_graphed_callables.<locals>.make_graphed_autograd_function.<locals>.__annotate__  s     	R 	Rv 	R& 	Rr   c                    < \         P                  P                  P                  ! V !  pSP                  ! \        V4      S,           !  p\         P                  P                  P                  VS4      # r,   )r$   utils_pytreearg_tree_leavesapplyr   tree_unflatten)r   flatten_user_argsoutr   r   r   s   *  r   functionalizedVmake_graphed_callables.<locals>.make_graphed_autograd_function.<locals>.functionalized  sU     % 3 3 C CY O--%(9":]"JLC;;&&55c;PQQr   )r$   r   Function)r   r   r   r   r   r   r   r   r   r   r   s   fffffffff @r   make_graphed_autograd_function>make_graphed_callables.<locals>.make_graphed_autograd_function  s4    	 	enn-- 	B	R 	R r   c          
     ,    V ^8  d   QhRRRRRRRRRR/# )	r   funcztorch.nn.Modulegraph_training_stater   graphedzCallable[_P, _R]orig_fwdr   r   )r   s   "r   r   r     s:      %&* * +	
 "r   c                &   a aaa R  V VVV3R llpV# )c               $    V ^8  d   QhRRRRRR/# )r   r   z_P.argsuser_kwargsz	_P.kwargsr   r   r   )r   s   "r   r   Jmake_graphed_callables.<locals>.make_graphed_forward.<locals>.__annotate__  s&     C C C	 Cb Cr   c                 F   < SP                   S8X  d	   S! V / VB # S! V / VB # r,   )training)r   r   r   r   r   r   s   *,r   new_fwdEmake_graphed_callables.<locals>.make_graphed_forward.<locals>.new_fwd  s0    }}(<<&	A[AA'BkBBr   r   )r   r   r   r   r   s   ffff r   make_graphed_forward4make_graphed_callables.<locals>.make_graphed_forward  s    C C r   zModules must not have hooks registered at the time they are passed. However, registering hooks on modules after passing them through make_graphed_callables is allowed.zIn any :class:`~torch.nn.Module` passed to :func:`~make_graphed_callables`, only parameters may be trainable. All buffers must have ``requires_grad=False``.zfIn the beta API, sample_args for each callable must contain only Tensors. Other types are not allowed.r   )*r$   is_autocast_enabledis_autocast_cache_enabledr   r   r   typingcastr   r   nnModuler   _backward_hooks_forward_hooks_forward_pre_hooksallbuffersr   r   r   append	TypeError
parametersr   r%   r   r&   r{   rp   rs   tree_leavesr   r   r   tree_flattenreversedr   reverse	enumerater   r   )+r   r   r   r   r3   just_one_callable_sample_argsflatten_sample_argscr   flatten_argper_callable_len_user_argsper_callable_module_paramsr   "per_callable_static_input_surfaces_
fwd_graphs
bwd_graphsmempoolr   r   grad_inputsr   outputs_gradvper_callable_static_outputs"per_callable_output_unflatten_specr   func_outputsflatten_outputsspec per_callable_static_grad_outputsper_callable_static_grad_inputsr   r   r   r   grad_idxr   r   retr   r   s+   &&&&&                                      r   r   r      sk   N   ""u'F'F'H'Hm
 	
  i'' L	E&#+$6DF{{5vs{);S)@#A;Oy/a))A%%&!+(()Q.,,-2"a  3EE333EEEE"1 
 kk))994@""5#56sHKHsssHKHHH^ ) 06 9L!L8K#d)8K!L "A ",Auxx!?!?allnRG  " s9~&*&A 	!;!>>>& ' *
 16c)n0EF0E1%))$$&0EJF05c)n0EF0E1%))$$&0EJF%)\!tG 
II			%))**,	-	-03%G1
,D$, 26tT,K+,++--99$+F$u$K$Kuu$K$KK|$q("'.."5"5 ,$u %';%uu %';%   &+U +9@+UU +9@+ & %)%7 #6 
#K	 - |[9 :'1
 
.. 
II #%)+&!$Yj!IdIYY__YW_55;L 6 !& 3 3 @ @ N#**5+AB*11$7 "J (*$&(#;>34,-<7ni
 $e $
FT$
ee $
FT$
 
 uJJuuJJJ|q 99#nn11( 5 T,@ T55 T,@ TT!&&W2E&W&W2E&W!W $!3 2  :  'C   [%<"))+h*?@A"))$/ ( ##56(//0CD'../ABA<F %,,.#++-2h $&CY'40qMqM&)&).q1.1'*,Q/+A.

 dEHHOO,, 0dmmWdllDL JJtJJwE (H 1v:w "M"*
 GF 
.	-	-< 655, :99s]   b4Ab9b>;'c:'cBc($c*c8>cc!Ac5(c5c	!c25d)r    r&   r   r   r   )   FN)'__conditional_annotations__
__future__r   r   collections.abcr   r   r   r   typing_extensionsr   r	   r
   r$   r   	torch.xpur   _utilsr   __all__r   r   hasattr_C__dict__torch._Cr   r   r   r    r&   r   r   r  r  r   r   r   r   )r*  s   @r   <module>r5     s1   " "  $ 5 5 6 6   &   T]t_uxx)**%0%=EHHk"2=>V2WEHH./9D':EHH56 V U+<
^,y ^,B3( 3(l  %xx#v+1FF 9 F 
 
 
( 
(n nr   