
    7h                    ^   S SK Jr  S SKrS SKrS SKrS SKrS SKrS SKrS SKJ	r	  S SK
Jr  S SKJrJr  S SKrS SKJr  S SKJr  S SKJr  S	S
KJrJr  S	SKJr  S	SKJrJrJr  S	SKJ r J!r!J"r"J#r#J$r$J%r%J&r&  S	SK'J(r(  \RR                  " \*5      r+\RX                  R[                  \*S5      r.\(       a  S SK/J0r0  S"S jr1S"S jr2    S"S jr3    S"S jr4\ " S S5      5       r5    S#S jr6          S$S jr7    S%S jr8S&S jr9S r:S r;    S"S jr<S'S jr=S(S jr>S  r?        S)S! jr@g)*    )annotationsN)defaultdict)	dataclass)AnyTYPE_CHECKING)trace_structured)StorageWeakRef)
OrderedSet   )configir)WeakDep)estimate_peak_memoryFreeableInputBufferget_freeable_input_buf)contains_collectivecontains_waitfind_recursive_deps_of_nodefind_recursive_users_of_nodeis_collectiveis_fallback_opis_wait)Voverlap)BaseSchedulerNodec                    [        U SSSS9$ )z/
Greedily schedules waits as late as possible.
FTraise_comms
sink_waitsreorder_for_overlap_schedule_for_commsnodess    O/var/www/fran/franai/venv/lib/python3.13/site-packages/torch/_inductor/comms.pyr   r   *   s     Ed     c                    [        U SSSS9$ )z0
Greedily schedules comms as early as possible.
TFr   r!   r#   s    r%   r   r   3   s     DU r&   c                    [        U SSSS9$ )a  
This achieves the following overall scheduling procedure:
    Step 1: Given that we've currently scheduled comm N, we now schedule all compute nodes
        that are required for comm N + 1 but do not depend on comm N, to run at the same time with comm N.
    Step 2: If all those compute nodes are sufficient to overlap comm N, we're done.
        Otherwise, we now need to look elsewhere to find compute that overlaps with comm N.
        We prioritize compute nodes that are needed sooner.
    Step 3: We schedule the compute nodes dependent on comm N and required for comm N + 1.
    Step 4: We schedule comm N + 1.
    Repeat this for subsequent comm nodes.
Tr   r!   r#   s    r%   reorder_compute_for_overlapr)   <   s     DTt r&   c           
       ^ [        U 5      u  pU Vs0 s H  o3X#   R                  _M     nn[        U Vs/ s H  o4U   PM	     sn5      n[        U Vs/ s H  o2U   R                  PM     sn5      nSU SU S3m/ SQnUR	                  5        VVs/ s HH  u  p8[        U5      UR                  UR                  UR                  UR                  UR                  /PMJ     n	nn[        R                  R                  S5      (       a  SSKJn
  TU
" U	US9-  mO8TS	-  mT[        U5      S
-   -  mTS
R                  [        [        U	5      5      -  m[         R#                  T5        [%        SS U4S jS9  U$ s  snf s  snf s  snf s  snnf )a+  
Reorders communication ops relative to computation ops to improve communication-compute overlapping and hide comm
latency.  Stops moving a particular op if it reaches a point that would have increased the peak memory footprint.

Currently, follows these heuristics (subject to change or tune):
- never reorders collectives relative to one another, for SPMD safety
- has an option for per-collective prefetch limit, but does not enable it by default
- limits the total number of reorder steps to some factor of the graph size to prevent worst-case quadratic
  performance

Prerequisite: sink_comms_and_waits - ensure comm and wait nodes are scheduled as late as possible, respecting data
dependencies.  That allows reorder_communication_preserving_peak_memory to take a best case peak-memory snapshot,
and then monotonically improve latency by moving collectives backward in time.

Peak memory impact is computed in an iterative fashion.  First, memory use at each timestep is computed, and global
peak memory is computed as a max over timesteps.  Then, when swapping any two adjacent nodes, only the curr-memory
for the earlier of the nodes after the swap is affected.  This enables checking step by step whether a swap is
peak-memory-safe, and bailing out if not.  Example:

0   n0      C0
1   n1      C0 + Allocs(n1) - Frees(n1)
2   n2      C0 + Allocs(n1) - Frees(n1) + Allocs(n2) - Frees(n2)

0   n0      C0
1   n2      C0 + Allocs(n2) - Frees(n2)    <-- After moving n2 to Time 1, only time1 memory changes
2   n1      C0 + Allocs(n2) - Frees(n2) + Allocs(n1) - Frees(n1)

zAreorder_communication_preserving_peak_memory improved overlap by z
 ns after z reorders.
)zCollective nodezinitial exposedzfinal exposedimprovementzlimiting factormovestabulater   )r-   )headersz>Please `pip install tabulate` to nicely render overlap stats.

artifactc                     SSS.$ )N,reorder_communication_preserving_peak_memorystring)nameencoding r6   r&   r%   <lambda>>reorder_communication_preserving_peak_memory.<locals>.<lambda>   s    B 
r&   c                    > T $ Nr6   )reorder_log_strs   r%   r7   r8      s    ?r&   )metadata_fn
payload_fn)6_reorder_communication_preserving_peak_memory_internalr+   sumr,   itemsnode_summaryinitial_exposedfinal_exposedlimiting_factor	importlibutil	find_specr-   strjoinmapoverlap_loginfor   )r$   reordered_snodes
node_statssnoder+   total_improvementtotal_movesr.   node_reorder_inforowsr-   r;   s              @r%   r2   r2   O   s   @ 	?vF ! FPPZE*+777ZKP[I[E/[IJJGJ5%(..JGHK LL]K^ _l	, G" )3(8(8(:
 );$E --++))--##	
 ); 	 
 ~~
++%8
 	

 	M	
 	3w<$..499Sd^44_%
 + e QIG
s   E9E>FAFc                  b    \ rS rSr% SrSrS\S'   SrS\S'   SrS\S	'   S
r	S\S'   \
S 5       rSrg)ReorderInfo   z=
Debug info describing how an individual snode was reordered
floatrB   rC   NonerH   rD   r   intr,   c                4    U R                   U R                  -
  $ r:   )rB   rC   )selfs    r%   r+   ReorderInfo.improvement   s    ##d&8&888r&   r6   N)__name__
__module____qualname____firstlineno____doc__rB   __annotations__rC   rD   r,   propertyr+   __static_attributes__r6   r&   r%   rU   rU      sB      OUM5!OS!E3N9 9r&   rU   c                P  ^^ [        U 5      S-  nSn[        U 5      n[        R                  b  [        R                  n[        [        R
                  R                  R                  5       5      n[        [        R
                  R                  5       5      n[        X5      n[        XU5      u  pxU  V	s0 s H  o[        U	5      _M     sn	m0 n
U4S jn[        U 5       GH  u  p[        U	5      (       d  M  [        5       =oU	'   U" XUS-   S 5      =Ul        Ul        X!:  a	  SUl        MN  [%        US-
  SS5       GHT  nX   nU['        SX-
  5      :  a
  SUl          M  [        U5      (       a
  S	Ul          M  [        U	R(                   Vs/ s H  nUR*                  PM     sn5      m[-        U4S
 jUR/                  5        5       5      (       a  [1        U5      (       d  SUl          GM  XxU   -
  XS-
     X   -
  :  a  SUl          GM/  UR                   TU	   :  a  SUl          GMM  U=R2                  S-  sl        US-  nX   nXS-      X'   UXS-   '   XS-      X   -
  nX   XS-
     -
  nX   U-
  U-   X'   U" XUS-   S 5      Ul        GMW     GM     X
4$ s  sn	f s  snf )z|
Internal testing helper that also returns debug info.
Returns:
    - reordered snodes list
    - dict {snode: ReorderInfo}
d   r   Nc                   > [        U 5      nSnU H/  n[        U5      (       a  M  [        U5      (       a    OUTU   -  nM1     [        SX#-
  5      $ )N        r   )estimate_op_runtimer   r   max)collective_snoderemaining_snodes	comm_timecompute_timerO   runtimess        r%   exposed_communication_timeZ_reorder_communication_preserving_peak_memory_internal.<locals>.exposed_communication_time   s[    '(89	%E"5))U## HUO+L & 1i.//r&   r   z
move limitrW   zprefetch limitzcollective orderingc              3  H   >#    U  H  oR                  5       T;   v   M     g 7fr:   )get_name).0o	dep_namess     r%   	<genexpr>I_reorder_communication_preserving_peak_memory_internal.<locals>.<genexpr>   s      7O!JJLI-7Os   "zdata dependencyzpeak memoryzsufficient overlapping)lenr   reorder_prefetch_limitr
   r   graphgraph_inputskeysget_output_namesr   r   rj   	enumerater   rU   rB   rC   rD   rangerk   unmet_dependenciesr4   anyget_outputsr   r,   )r$   
MOVE_LIMITrQ   PER_COLLECTIVE_PREFETCH_LIMITr}   graph_outputsname_to_freeable_input_bufpeak_memorycurr_memoryrO   statsrq   ireorder_infoj
prev_snodestmpj_plus_one_allocj_allocrw   rp   s                       @@r%   r>   r>      s    Vs"JK$'K!$$0(.(E(E%$.qww/C/C/H/H/J$KL%/0H0H0J%KMAWB  4M K @FFve*511vFH 35E0 f%u%%*5-7L<*5Q/BL(<+E (/;,1q5"b)#Y
s1a?@@3CL0&z223HL0&8P8P'Q8P18P'QR	 7A7M7M7O  '
333DL0Q/+!e2D{~2UU3@L0--?3KL0""a'"q i"q5M	 #1u#.1u#5#F %.;1u+==!,'!9<L!L-G!a%'?.*; * &T =} GJ (Rs   4JJ#c                  ^^^^^^^^^^^^ 0 n0 m0 0 0 smmm[        U 5       H|  u  pVUR                  5        H  nXdU'   M	     UR                  5        H  nUTU'   M
     UTUR                  5       '   UR                  5       n	[        R
                  TU	'   STU	'   UTU	'   M~     Sn
U  H  nU(       ab  [        U5      (       aR  U
TUR                  5       '   UR                   H(  nTU   R                  5       n[        TU   U
5      TU'   M*     U
S-  n
Ml  U(       d  Mu  [        U5      (       d  M  STUR                  5       '   M      " UUUU4S jS5      mU  Vs0 s H   nU[        S UR                   5       5      _M"     snm/ m[        [        5      mU  Vs0 s H  of[        U5      _M     snmTR                  5        HN  u  pm[        U5      S:X  a  [         R"                  " TT" U5      5        U H  nTU   R%                  U5        M     MP     / mUUUUU4S jmU4S jmUUUU4S jn[        T5      (       aZ  [         R&                  " T5      R(                  nU(       a  [        U5      (       a	  U" U5        OT" U5        [        T5      (       a  MZ  TR                  5        H  u  pm[        U5      S:X  a  M   S	T 35       e   T$ s  snf s  snf )
a  
Schedule `snodes` for various comm optimization objectives.

Args:
    snodes: the nodes to be scheduled.
    raise_comms: whether to greedily schedule collectives as early as possible
    sink_wait: whether to greedily schedule waits as late as possible
    reorder_compute_for_overlap: whether to reorder compute nodes to
        optimize for compute/communication overlapping.

Returns:
    The new schedule order.

Some notes on the synergy between different options:
    - `raise_comms` provides more overlapping oppurtunies for `reorder_compute_for_overlap`.
    - When both `raise_comms` and `sink_waits` is `True`, `raise_comms` is prioritized.
r   r   c                  2   > \ rS rSrSU UUU4S jjrS rSrg)$_schedule_for_comm.<locals>.RunnableiP  c                   > Xl         [        [        UR                  5       5      5      nTU   R	                  5       nTU   TU   TU   4U l        g r:   )rO   nextiterget_operation_namesrt   score)r\   rO   r4   
fused_namename_to_fused_nodescores_0scores_1scores_2s       r%   __init__-_schedule_for_comm.<locals>.Runnable.__init__Q  sT    JU6689:D+D1::<J$$$DJr&   c                4    U R                   UR                   :  $ r:   r   )r\   others     r%   __lt__+_schedule_for_comm.<locals>.Runnable.__lt__[  s    ::++r&   )r   rO   N)returnrY   )r^   r_   r`   ra   r   r   re   )r   r   r   r   s   r%   Runnabler   P  s    	 		,r&   r   c              3  8   #    U  H  oR                   v   M     g 7fr:   )r4   )ru   deps     r%   rx   %_schedule_for_comm.<locals>.<genexpr>_  s     G.Fs((.Fs   c                   > TR                  U 5        U R                  5        HT  nTU    HH  n TU    R                  U5        [        TU    5      S:X  d  M+  [        R
                  " TT" U 5      5        MJ     MV     g)zE
Schedules `snode` and put all unblocked nodes onto the ready queue.
r   N)appendget_buffer_namesremoverz   heapqheappush)rO   buf_namer   buffer_usersready	scheduled
unmet_depss     r%   schedule$_schedule_for_comm.<locals>.scheduleo  sl     	..0H%h/5!((2z%()Q.NN5(5/: 0 1r&   c                    > T V s/ s H=  n [        U R                  5      (       a  M  [        U R                  5      (       a  M;  U PM?     nn [        U5      S:X  a  g[	        US S9$ s  sn f )zP
Return the next node in the ready queue that's neither a collective or
a wait.
r   Nc                    U R                   $ r:   r   xs    r%   r7   G_schedule_for_comm.<locals>.get_overlapping_candidate.<locals>.<lambda>  s    QWWr&   key)r   rO   r   rz   min)r   
candidatesr   s     r%   get_overlapping_candidate5_schedule_for_comm.<locals>.get_overlapping_candidatez  sf     
&qww/ 8Eagg8N  	 

 z?a:#455
s   A'A'A'c                  > [        U 5      (       d   eT" U 5        TU    nUS:  aQ  T" 5       =nbG  TR                  U5        T" UR                  5        UTUR                     -  nUS:  a  T" 5       =nb  MG  [        R                  " T5        g)z
Schedules collective node `snode`, along with one or more compute nodes
to overlap with it. The strategy is described in the comment of
`reorder_compute_for_overlap`.
r   N)r   r   rO   r   heapify)rO   collective_cost	candidater   r   r   snode_to_costs      r%   schedule_collective_for_overlap;_schedule_for_comm.<locals>.schedule_collective_for_overlap  s     #5))))'.a799FLL#Y__%}Y__==O a799F
 	er&   z;Detected unscheduled nodes. Nodes with unmet dependencies: )r   r   r   rt   sysmaxsizer   	ancestorsr   r   r
   r   r   rj   r@   rz   r   r   addheappoprO   )r$   r   r   r    buf_name_to_snodeidxrO   r   op_name	node_namecomm_idxancestoranc_fused_namedepsr   r   r   r   r   r   r   r   r   r   r   r   r   r   s                   @@@@@@@@@@@@r%   r"   r"     s   L #%r2 Hh'
..0H*/h' 1 002G*/w' 3/45>>+,NN$	!kk! ( H.u55)1HU^^%&!OO!3H!=!F!F!H+.x/G+R( , MHZM%00)*HU^^%& , ,  <E 	zGe.F.FGGG<J
 E=H=TLDJKF5/66FKM!'')t9>NN5(5/2C!!%(  * I	; 	;6 & e**e$**#6u#=#=+E2UO e** "'')4yA~ 	
I*V	
~ * Q< Ls   'KKc           	        [         R                  R                  5       (       d  U $ U  Vs/ s H  n[        U5      (       d  M  UPM     nn[	        S[        U5      5       H]  n[        [        XE   R                  5       5      5      nXES-
     R                  5        H  nXE   R                  [        XvS95        M     M_     U $ s  snf )z
Decide global ordering of comms, by just enforcing the ordering that's in the input graph
(might not be the same ordering as the eager mode program).
TODO: Come up with a better approach
r   mutating_buf)torchdistributedis_availabler   r   rz   r   r   r   add_fake_depr   )nodesname_to_bufr   n
comm_nodesr   r   bufs           r%   decide_global_ordering_of_commsr     s     ))++"=U&9!&<!UJ=1c*o&D!?!?!ABC!e$557CM&&ws'NO 8 ' L >s   CCc                    [         R                  S:X  a  U R                  5       nU$ [        [         R                  5      (       d   e[         R                  " U 5      nU$ )z2
Returns estimated op runtime in nanoseconds (ns)
default)r   rj   get_estimated_runtimecallable)rO   runtimes     r%   rj   rj     sU     !!Y.--/ N 223333,,U3Nr&   c           
        U R                  5       n[        U5      S:X  GaA  Sn[        U R                  [        R
                  [        R                  45      (       a  SU R                  R                   S3nU R                  5        Vs/ s H  o3R                  R                  5       PM     nnSR                  U Vs/ s HA  n[        U[        R                  5      (       a  SUR                   SUR                   S3OSPMC     sn5      n U R                  R                  5       nU R                  R                  R                    U U SU SU R#                  5       S S	3$ / nU H  n	UR%                  ['        U	5      5        M     U R                  R                    S
SR                  U5       3$ s  snf s  snf ! [         a    Sn Nf = f)Nr    z (),z (size=z	, stride=z.0fz ns): z, )	get_nodesrz   
isinstancenoder   ExternKernelOut_CollectiveKernelpython_kernel_nameget_output_specrI   Layoutsizestridemaybe_get_nameAttributeError	__class__r^   r   r   rA   )
rO   r$   detailchildlayoutslayoutout_tensor_infor   	summarieschild_snodes
             r%   rA   rA     s   __F
6{aejj2#5#5r7K7K"LMM%**778:F=B__=NO=NE::--/=NO((
 &	 &F fbii00 &++ia@ &	
	

113I **&&//08II;VXY^YtYtYvwzX{{  A  	A Ik23 oo&&'r$))I*>)?@@) P  	I	s   #F1<AF6F; ;G
	G
c                X   SnS nS n[        U 5       H  u  pEUci  [        U5      (       a  U[        U5      -  nUR                  nO)[	        UR                  5      (       a  OU[        U5      -  nU" U[        U5       5        Mq  [        U5      (       a/  U[        U5      -  nUR                  nU" U[        U5       5        M  [	        UR                  5      (       a  U" U[        U5       5        S nM  U" US[        U5       35        M     [        R                  SUS-  S-   35        g )Nri   c                :    [         R                  U S SU 35        g )Nz>6r   )rK   debug)stepmsgs     r%   step_log#visualize_overlap.<locals>.step_log  s    T"IRu-.r&   z| zEst. runtime (ms): i  )r   r   rj   r   r   rA   rK   r  )ordertotal_est_runtimecur_comm_noder	  r  rO   s         r%   visualize_overlapr    s#     #M/ !' "5))!%8%??! %

$$ !%8%??!Tl5124"5))!%8%??! %

,u"5!68$$,u"5!68 $L$7#89:- (. 
/$6=>?r&   c                @   U n[        [        R                  R                  R	                  5       5      n[        [        R                  R                  5       5      n[        R                   GHX  n[        U[        5      (       a  U[        5       ;   a  [        5       U   n[        U5      (       d   SU S35       e[        U [        X5      U5      u  pV[        R                  R!                  5       S:X  a)  ["        R%                  SU SU< S35         ['        U5        [*        R*                  " 5       nU" U5      n[*        R*                  " 5       U-
  n	[        R                  R!                  5       S:X  a(  ["        R%                  S	U S
U	 S35         ['        U5        [        U [        X5      U5      u  pV[-        SU< 35        GM[     U$ ! [(         a  n["        R%                  SUS9   S nANS nAff = f! [(         a  n["        R%                  SUS9   S nAN|S nAff = f)Nz3Invalid reorder_compute_and_comm_for_overlap pass: z is not callabler   z.==== Visualize overlap before reordering pass z, peak_memory=z ====r   )exc_infoz-==== Visualize overlap after reordering pass z	 (ran in z	 sec)====zfinal peak_memory=)r
   r   r|   r}   r~   r   r   'reorder_for_compute_comm_overlap_passesr   rH   globalsr   r   r   r   r   get_rankrK   r  r  	Exceptiontimeprint)
r$   r  r}   r   pr   _et0ts
             r%   $reorder_compute_and_comm_for_overlapr    s    E$.qww/C/C/H/H/J$KL%/0H0H0J%KM;;a!wy.	!A{{ 	
A!DTU	
{ .*6@-
 %%'1,@?k^SXY2!%( YY[%IIK"%%'1,?s)A3iX2!%( .*6@-
 	#{n%&? <@ L#  2!!"q!12  2!!"q!12s0   G
G5

G2G--G25
H?HHc           
     ^	  ^^^^^^ [        U R                  5      m[        [         5      m[        [         5      m[        T5       H  u  pUR                  S:X  d  M  UR
                  [        R                  R                  R                  R                  :X  d  MU  UR                  S   R                  S:X  d   SU SUR                  S    S35       eUR                  S   nUR                  S   nUS:  a  TU   R                  U5        M  TU   R                  U5        M     UUU4S jn[        [         5      n[        T5       H  u  pUR                  S:X  d  M  UR
                  [        R                  R                  R                  R                  :X  d  MU  UnUR                  S   mTR                  S:X  d   S	T S
U  S35       eU" T5      (       d  M  UT   R                  U5        M     S nS mT H  nUR                  S:X  d  M  [        UR
                  [        R                   R"                  5      (       d  MJ  UR
                  R$                  R&                  (       d  Mq  U" U5      (       a  M  T" X&R)                  5       5      (       d  M   SU S35       e   UR+                  5        GH3  u  mn	[        U	5       GH  u  pTU   nUR                  S   TL d   eUR                  u  nmUS-   nU
[-        U	5      S-
  :  a  XS-      O[-        T5      S-
  nTX n[/        UU4S jU 5       5      (       a   ST SU SU  S35       eU H  nUR                  S:X  d  M  TUR                  ;   d  M'  UR
                  [        R                  R                  R                  R                  :w  d  Me  [1        UU4S jUR                   5       5      nUUl        M     GM     GM6     UR+                  5        H0  u  mn	[        U	5       H  u  pTU   nU R3                  U5        M     M2     T Hy  nUR                  S:X  d  M  UR
                  [        R                  R                  R                  R                  :X  d  MS  UR                  S   U;   d  Mh  U R3                  U5        M{     g)ab  
This FX graph pass replaces uses of FSDP2 unsharded params with their corresponding
graph intermediates that were fsdp.copy_ into the unsharded params in the original graph.

NOTE: Can only apply this pass to any of the FSDP2 unsharded params that have this pattern
(or repetition of): `resize_(full) -> copy_ -> resize_(0)`. Because of this, for partial-graph case
where `resize_(full) -> copy_` is in one graph and `resize_(0)` is in another graph, we can't
remove these resize and copy ops and thus we will have worse performance there.

In other words, "do we try to remove all the resize_(full) -> copy_ -> resize_(0) nodes for this unsharded param"
is actually a per-unsharded-param decision, since for each unsharded param, we look at its resize sequence pattern
(in `check_resize_pattern()`) to determine if its set of resize and copy nodes can be removed.
call_functionr   placeholderz1Resize can only operate on graph inputs, but got z# which is resizing non-graph-input r/   r   c                n  > TR                  U / 5      nTR                  U / 5      n[        U5      [        U5      :X  d2  [        R                  SU  S[        U5       S[        U5       S35        g[	        X5       H7  u  p4X4:  d  M  [        R                  SU  STU    SU S	TU    SU S
35          g   g)NzH
Unequal number of resize-to-full and resize-to-0 nodes for graph input z:
z vs. zK.
Skipping `remove_fsdp2_unsharded_param_graph_input_usage` FX graph pass.
Fz
For graph input z: resize-to-full node z
 at index z 
happens after resize-to-0 node zd.
Skipping `remove_fsdp2_unsharded_param_graph_input_usage` FX graph pass for that unsharded param.
T)getrz   logwarningzip)graph_inputresized_to_full_idxesresized_to_0_idxesresize_to_full_idxresize_to_0_idx&graph_input_to_resized_to_0_node_idxes)graph_input_to_resized_to_full_node_idxes	node_lists        r%   check_resize_patternLremove_fsdp2_unsharded_param_graph_input_usage.<locals>.check_resize_patternY  s    !J M M!
 DGGUWX()S1C-DDKKHHS} U E#&8"9!: ;  47!4
/ "43I>P4Q3RR\]o\p q  )/ :;:oEV W 4
 r&   z\
Assumed all FSDP2 `unsharded_param`s to be graph input, but it's not true!
Offending node: z	. Graph: c                    U R                   [        R                  R                  R                  R
                  :H  =(       d;    U R                   [        R                  R                  R                  R
                  :H  $ r:   )targetr   opsfsdpcopy_r   inductorresize_storage_bytes_)r   s    r%   is_allowed_mutationKremove_fsdp2_unsharded_param_graph_input_usage.<locals>.is_allowed_mutation  sO    KK599>>//777 O{{eii00FFNNN	
r&   c           	        [        U R                  [        R                  R                  5      (       aj  [        U R                  R                  R                  5       VVs/ s H3  u  p#UR                  c  M  UR                  R                  (       d  M1  UPM5     snnO/ n[        U Vs/ s H6  n[        U R                  U   R                  S   R                  5       5      PM8     sn5      n[        U Vs/ s H)  n[        UR                  S   R                  5       5      PM+     sn5      n[        XW-  5      S:  $ s  snnf s  snf s  snf )Nvalr   )r   r0  r   _ops
OpOverloadr   _schema	arguments
alias_infois_writer
   r	   argsmetauntyped_storagerz   )r   unsharded_paramsr   r   mutated_arg_idxesmutated_node_arg_storagesunsharded_paramstorages_of_unsharded_paramss           r%   -is_node_mutating_unsharded_param_or_its_aliaseremove_fsdp2_unsharded_param_graph_input_usage.<locals>.is_node_mutating_unsharded_param_or_its_alias  s)    $++uzz'<'<== &dkk&9&9&C&CDDDA<< 010E0E D  	 %/ +*A tyy|007GGIJ*%
! (2 (8'7O 33E:JJLM'7(
$ ,KLqPP)s    D=7D=D=)=E60EzdUser mutation on FSDP2 unsharded param is not allowed when Traceable FSDP2 is used. Violating node: c              3  8   >#    U  H  nT" UT/5      v   M     g 7fr:   r6   )ru   r   rH  rF  s     r%   rx   Aremove_fsdp2_unsharded_param_graph_input_usage.<locals>.<genexpr>  s%      *D >d_DUVV*s   z(Assumed no ops mutating unsharded param z in subgraph z, but it's not true!
Graph: c              3  6   >#    U  H  nUTL a  TOUv   M     g 7fr:   r6   )ru   argreplacementrF  s     r%   rx   rK    s$      %#,C (+o'=3F#,s   N)listr   r   r   opr0  r   r1  r4  r5  r   r@  r   r2  r3  r   r:  r;  r<  
is_mutabler~   r@   rz   r   tuple
erase_node)r|   r   r   r%  new_sizer-  'unsharded_param_to_fsdp_copy_node_idxesfsdp_copy_noder6  fsdp_copy_node_idxesr   fsdp_copy_node_idxr  subgraph_start_idxsubgraph_end_idxsubgraph_nodesnew_argsr*  r+  rH  r,  rN  rF  s                    @@@@@@r%   .remove_fsdp2_unsharded_param_graph_input_usager]  7  sb    U[[!I 1<D0A--8->*y)	GG&uyy11GGOOO99Q<??m3  :2267Z[_[d[def[gZh i6 3 ))A,Kyy|H!|9+FMMcR6{CJJ3O *"J /:$.?+y)	77o%$++9M9M9U9U*U!N"iilO"%%6  = !5' 29 6 $O447HOOPST *
Q4 GG&4;;

(=(=>>##...'--DBBD  eeidj k  : 
1	6	6	8	%./C%D!A&'9:N!&&q)_<<<+00NA{!3a!7 s/0144 %U+^a' 
 ''9KN *   ))8(9~FV Ww   'GG.'4994uyy'9'9'O'O'W'WW$ %#'99%  H !)DI ') &E 
9J 
1	6	6	8	%./C%D!A&'9:N^, &E 
9 GG&uyy11GGOOO		! GGT" r&   c                @  ^	  SS K m	T	R                  R                  5       (       d   eT	R                  R                  R
                  (       a%  T	R                  R                  R                  (       d   e SSK
JnJnJnJnJn   U	4S jnU" 5       nU" U" T	R                  R                  R
                  R                   U" ["        R$                  U" T	R                  R&                  R(                  R                   U" S5      U" S5      U" S5      U" S5      U" S	5      U" S
5      U" S5      U" S5      U" S5      5
      U" S5      5      U" S5      U" S5      5      US S9SU	4S jj5       nU" U 5        UR+                  U 5        g ! [        [        [        4 a     g f = f)Nr   r   )CallFunction
KeywordArgMatchPatternMatcherPassregister_graph_patternc                Z  > [        U R                  5      nU H  nUR                  [        R                  :X  d  M#  UR
                  S   R                  TR                  R                  R                  R                  L d  Mi  UR
                  S   S:X  d  M~  U R                  U5        M     g )Nr   r   )rO  r   r0  operatorgetitemr@  r1  r2  all_gather_copy_inr   rS  )gr,  r   r   s      r%   remove_unused_getitem8reinplace_fsdp_all_gather.<locals>.remove_unused_getitem   sp    M	AH,,,FF1I$$		(I(I(Q(QQFF1INQ r&   all_gather_inputsinp_split_sizesall_gather_input_numel
world_sizerankdtypedevicegroup_name_inner"allocate_memory_from_process_groupitem_idx
group_size
group_namec                &    U R                   S   S:H  $ )Nrt  r   )kwargs)matchs    r%   r7   +reinplace_fsdp_all_gather.<locals>.<lambda>D  s    %,,z":a"?r&   )	pass_dictextra_checkc                   > U4S jnU R                  UUS   US   US   US   US   US   US   US	   US
   US   US   /5        g )Nc                    > U S S nU S   nU S   nTR                   R                  R                  R                  " U6 nUS   nUS   nTR                   R                  R
                  R                  XRX6S9nU$ )NrW   r   r   )out)r1  r2  rg  r   _c10d_functionalall_gather_into_tensor_out)	r@  copy_in_argsru  rv  rg  rf  	getitem_1all_gather_into_tensorr   s	           r%   replEreinplace_fsdp_all_gather.<locals>.reinplace_all_gather.<locals>.replG  s      9LbJbJ!&!B!B!J!J" )+G*1-I		**EEMM N  #
 *)r&   rk  rl  rm  rn  ro  rp  rq  rr  rs  ru  rv  )replace_by_example)ry  r@  rx  r  r   s       r%   reinplace_all_gather7reinplace_fsdp_all_gather.<locals>.reinplace_all_gather-  s{    4	*$ 	  *+()/0|$vwx )*;<|$|$	
r&   )ry  ra  )5torch.distributed.fsdp._fully_shard._fsdp_collectivesr   r   r1  r  r  r  ImportErrorr   AssertionErrorpattern_matcherr_  r`  ra  rb  rc  r   re  rf  r2  rg  apply)
r|   r_  r`  ra  rb  rc  ri  
graph_passr  r   s
            @r%   reinplace_fsdp_all_gatherr    s   
D  --//// II&&==		**EE	
FE
  	  $%JII&&==EE  IINN55==230178|,v&w'x(12CD :&  |$|$'	
* ?/2"
32"
H % UM 8 s   A1F FFc                    [        U [        R                  R                  R                  [        R                  R                  R
                  45      (       a   e[        U R                  5       SS  5      $ )N   )r   r   	_inductor	schedulerFusedSchedulerNodeGroupedSchedulerNoderZ   rt   )rO   s    r%   
get_op_idxr  n  sb    OO%%88OO%%::	
    u~~#$$r&   c           	     	  ^^^ ^! SSK Jm   / n[        [           " 5       nSnSn0 n0 n0 m!U U!4S jn	U  GH"  n
[	        U
R
                  [        R                  R                  R                  R                  S9(       Ga  [        U4S jU
R                   5       5      (       Ga  SnU
n[        5       n[        UUUT5        [        [        R                  R                  R                  R                  [        R                  R                  R                  R                  [        R                  R                  R                   R                  /5      m[#        UUUTUU 4S jS	9  [%        US
 S9n['        U5      nSn[)        ['        U5      5       H^  nUU   n[+        UR
                  [        R                  R                  R                   R                  5      (       a  US-  nUS:  d  M\  Un  O   US U nS n[)        ['        U5      S-
  5       H9  n[-        UUS-      R
                  [.        R0                  5      (       d  M4  US-   n  O   Uc   eU	" US U 5      nU	" UUS  5      nUUU'   GM;  [+        U
R
                  [        R                  R                  R2                  R                  5      (       d  GM  SnU
n[        5       n[#        UUUT5        [%        US S9nS n[)        ['        U5      S-
  5       H9  n[-        UUS-      R
                  [.        R0                  5      (       d  M4  US-   n  O   Uc   eU	" US U 5      nU	" UUS  5      nUUU'   GM%     ['        T!5      S:  d   eU(       a  ['        U5      S:  d   eU(       a  ['        U5      S:  d   eU  HS  n
U
R5                  5       T!;   a  T!U
R5                  5          n
X;   a  M1  UR7                  U
5        UR9                  U
5        MU     S nUR;                  5        Hk  u  nnUb`  [=        [?        URA                  5       5      5      nURC                  5        H*  nURE                  [G        UR5                  5       US95        M,     UnMm     S nUR;                  5        Hk  u  nnUb`  [=        [?        URA                  5       5      5      nURC                  5        H*  nURE                  [G        UR5                  5       US95        M,     UnMm     U$ )Nr   )r  Fc                   > TR                   R                  U 5      nU  H  nUTUR                  5       '   M     UTUR                  5       '   U$ r:   )r  creatert   )snodes_to_group
group_noderO   r  snode_name_to_final_snodes      r%   _create_group_node:enforce_comm_ordering_for_fsdp.<locals>._create_group_node  sO    33::?K
$E:D%enn&67 %;E!*"5"5"78r&   )rP  c              3     >#    U  HJ  n[        TU   R                  [        R                  R                  R
                  R                  5      v   ML     g 7fr:   )r   r   r   r1  r2  rg  r   )ru   r   r   s     r%   rx   1enforce_comm_ordering_for_fsdp.<locals>.<genexpr>  sJ      
 % "1%**EIINN,M,M,U,U  %s   AATc                   > [        U TR                  5      =(       d6    [        U TR                  5      =(       a    U R                  R                  T;   (       + $ r:   )r   NopKernelSchedulerNodeExternKernelSchedulerNoder   op_overload)r   allowed_opsr  s    r%   r7   0enforce_comm_ordering_for_fsdp.<locals>.<lambda>  sD    q)"B"BC "1i&I&IJ >FF..+=	'r&   )criteria_cbc                    [        U 5      $ r:   r  r   s    r%   r7   r        JqMr&   r   r   c                    [        U 5      $ r:   r  r   s    r%   r7   r    r  r&   r   )$r   r  r
   r   r   r   r   r1  r  r  r   r   r   r   wait_tensorr2  split_with_sizes_copyr   sortedrz   r   r   r   r   _WaitKernel	chunk_catrt   r   r   r@   r   r   r   r   r   r   )"r$   r   r   	new_orderr   	ag_exists	rs_exists$ag_grouped_node_to_wait_grouped_node$rs_grouped_node_to_wait_grouped_noder  rO   ag_snodeag_related_snode_setag_related_snodesend_idx_of_current_ag_blockcopy_out_countr   	cur_snodewait_node_idxag_group_nodeag_wait_group_noders_snoders_related_snode_setrs_related_snodesrs_group_noders_wait_group_nodeprev_ag_waitwait_group_noder   rv   prev_rs_waitr  r  r  s"     `                            @@@r%   enforce_comm_ordering_for_fsdpr  y  s   
 )+I3!III+-(+-( " JJ59955PPXX
 
 
 __	
 
 
 IHLVL  ($"	 %II..IIQQII..::BBIINN88@@K )$" !'$*A! +..?*@'N3012-a0	!NNEIINN$H$H$P$P  #a'N!A%23/ 3 !22N3N O !M301A56/A6;;R^^LL$%EM 7 !,,,./@-/PQM "44Emn4U!VBT0? EJJ		(@(@(H(HIIIH MWL ($"	 !'$*A!
 !M301A56/A6;;R^^LL$%EM 7 !,,,./@-/PQM "44Emn4U!VBT0?] ` ()A---781<<<781<<< >>88-enn.>?Ee  L*N*T*T*V&#]%C%C%E FGL!--/**AJJL|D 0 ' +W L*N*T*T*V&#]%C%C%E FGL!--/**AJJL|D 0 ' +W r&   )r$   list[BaseSchedulerNode]r   r  )r$   r  r   zDtuple[list[BaseSchedulerNode], dict[BaseSchedulerNode, ReorderInfo]])
r$   r  r   boolr   r  r    r  r   r  )r   r  r   r  )rO   r   r   rX   )r|   torch.fx.Graph)r|   r  r   rY   )r$   1list[torch._inductor.scheduler.BaseSchedulerNode]r   z4dict[str, torch._inductor.scheduler.SchedulerBuffer]r   zdict[str, BaseSchedulerNode]r   r  )A
__future__r   r   rE   loggingre  r   r  collectionsr   dataclassesr   typingr   r   r   torch._loggingr    torch.multiprocessing.reductionsr	   torch.utils._ordered_setr
   r   r   r   dependenciesr   memoryr   r   r   utilsr   r   r   r   r   r   r   virtualizedr   	getLoggerr^   r"  _logginggetArtifactLoggerrK   torch._inductor.schedulerr   r   r   r)   r2   rU   r>   r"   r   rj   rA   r  r  r]  r  r  r  r6   r&   r%   <module>r     s   #     
  # ! %  + ; /  ! U U    !nn..xC;#&T#TTn 9 9 9V#VIVrW#WW W 	W
 Wt",	A:#L&#&&RA#Hpf%n=nEn 5n 7	nr&   