Ë
    g^(h”~  ã                  ó¶  — d dl m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	m
Z
 d dlZd dlmZ d dlmZ ddlmZmZ dd	lmZ dd
lmZmZmZmZmZmZmZ  ej:                  e«      Zej@                  jC                  ed«      Z"e
rddl#m$Z$ dd„Z%dd„Z&	 	 	 	 dd„Z'	 	 	 	 	 	 	 	 	 	 dd„Z(	 	 	 	 dd„Z)dd„Z*d„ Z+d„ Z,	 	 	 	 dd„Z-dd„Z.dd„Z/d„ Z0	 	 	 	 	 	 	 	 d d„Z1y)!é    )ÚannotationsN)Údefaultdict)ÚAnyÚTYPE_CHECKING)ÚStorageWeakRef)Ú
OrderedSeté   )ÚconfigÚir)ÚWeakDep)Úcontains_collectiveÚcontains_waitÚfind_recursive_deps_of_nodeÚfind_recursive_users_of_nodeÚis_collectiveÚis_fallback_opÚis_waitÚoverlap)ÚBaseSchedulerNodec                ó    — t        | ddd¬«      S )z7
    Greedily schedules waits as late as possible.
    FT©Úraise_commsÚ
sink_waitsÚreorder_for_overlap©Ú_schedule_for_comm©Úsnodess    úS/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/_inductor/comms.pyr   r   $   s   € ô Ø˜E¨dÈôð ó    c                ó    — t        | ddd¬«      S )z8
    Greedily schedules comms as early as possible.
    TFr   r   r   s    r   r   r   -   s   € ô Ø˜D¨UÈôð r    c                ó    — t        | ddd¬«      S )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#   6   s   € ô Ø˜D¨TÀtôð r    c                ó   ‡‡‡‡‡‡‡‡‡‡‡‡— i }i Ši i i cŠŠŠt        | «      D ]y  \  }}|j                  «       D ]  }|||<   Œ	 |j                  «       D ]  }|‰|<   Œ	 |‰|j                  «       <   |j                  «       }	t        j
                  ‰|	<   d‰|	<   |‰|	<   Œ{ d}
| D ]€  }|rZt        |«      rO|
‰|j                  «       <   |j                  D ]'  }‰|   j                  «       }t        ‰|   |
«      ‰|<   Œ) |
dz  }
Œ_|sŒbt        |«      sŒnd‰|j                  «       <   Œ‚  G ˆˆˆˆfd„d«      Š| D �ci c]  }|t        d„ |j                  D «       «      “Œ! c}Šg Št        t        «      Š| D �ci c]  }|t        |«      “Œ c}Š‰j                  «       D ]J  \  }}t        |«      dk(  rt!        j"                  ‰ ‰|«      «       |D ]  }‰|   j%                  |«       Œ ŒL g Šˆˆˆˆˆfd„Šˆfd„Šˆˆˆˆfd„}t        ‰«      rIt!        j&                  ‰«      j(                  }|rt        |«      r	 ||«       n ‰|«       t        ‰«      rŒI‰j                  «       D ]  \  }}t        |«      dk(  rŒJ d	‰› �«       ‚ ‰S c c}w c c}w )
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                  ó&   •— e Zd Zdˆ ˆˆˆfd„Zd„ Zy)ú$_schedule_for_comm.<locals>.Runnablec                ó¤   •— || _         t        t        |j                  «       «      «      }‰|   j	                  «       }‰|   ‰|   ‰|   f| _        y ©N)ÚsnodeÚnextÚiterÚget_operation_namesÚget_nameÚscore)Úselfr)   ÚnameÚ
fused_nameÚname_to_fused_nodeÚscores_0Úscores_1Úscores_2s       €€€€r   Ú__init__z-_schedule_for_comm.<locals>.Runnable.__init__‹   sV   ø€ ØˆDŒJÜœ˜U×6Ñ6Ó8Ó9Ó:ˆDØ+¨DÑ1×:Ñ:Ó<ˆJà˜Ñ$Ø˜Ñ$Ø˜Ñ$ðˆD�Jr    c                ó4   — | j                   |j                   k  S r(   ©r.   )r/   Úothers     r   Ú__lt__z+_schedule_for_comm.<locals>.Runnable.__lt__•   s   € Ø—:‘: §¡Ñ+Ð+r    N)ÚreturnÚNone)Ú__name__Ú
__module__Ú__qualname__r6   r:   )r2   r3   r4   r5   s   €€€€r   ÚRunnabler&   Š   s   ø„ ÷	ð 	ó	,r    r@   c              3  ó4   K  — | ]  }|j                   –— Œ y ­wr(   )r0   )Ú.0Údeps     r   ú	<genexpr>z%_schedule_for_comm.<locals>.<genexpr>™   s   è ø€ ÒG s˜#Ÿ(�(ÑGùs   ‚c                óê   •— ‰j                  | «       | j                  «       D ]N  }‰|   D ]D  } ‰|    j                  |«       t        ‰|    «      dk(  sŒ)t	        j
                  ‰ ‰| «      «       ŒF ŒP y)zU
        Schedules `snode` and put all unblocked nodes onto the ready queue.
        r   N)ÚappendÚget_buffer_namesÚremoveÚlenÚheapqÚheappush)r)   Úbuf_namer@   Úbuffer_usersÚreadyÚ	scheduledÚ
unmet_depss     €€€€€r   Úschedulez$_schedule_for_comm.<locals>.schedule©   sv   ø€ ð 	×Ñ˜ÔØ×.Ñ.Ó0ò 	;ˆHØ% hÑ/ò ;�Ø˜5Ñ!×(Ñ(¨Ô2Ü�z %Ñ(Ó)¨QÓ.Ü—N‘N 5©(°5«/Õ:ñ;ñ	;r    c                 óº   •— ‰D � cg c].  } t        | j                  «      st        | j                  «      s| ‘Œ0 }} t        |«      dk(  ryt	        |d„ ¬«      S c c} w )zh
        Return the next node in the ready queue that's neither a collective or
        a wait.
        r   Nc                ó   — | j                   S r(   r8   ©Úxs    r   ú<lambda>zG_schedule_for_comm.<locals>.get_overlapping_candidate.<locals>.<lambda>À   s
   € ¨Q¯W©W€ r    ©Úkey)r   r)   r   rI   Úmin)rU   Ú
candidatesrN   s     €r   Úget_overlapping_candidatez5_schedule_for_comm.<locals>.get_overlapping_candidate´   s]   ø€ ð ö
àÜ& q§w¡wÔ/¼ÀaÇgÁgÔ8Nò ð
ˆ
ð 
ô
 ˆz‹?˜aÒØÜ�:Ñ#4Ô5Ð5ùò
s   †3Ac                ó  •— t        | «      sJ ‚ ‰| «       ‰|    }|dkD  rM ‰«       x}�D‰j                  |«        ‰|j                  «       |‰|j                     z  }|dkD  r
 ‰«       x}�ŒDt        j                  ‰«       y)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   rH   r)   rJ   Úheapify)r)   Úcollective_costÚ	candidater[   rN   rQ   Úsnode_to_costs      €€€€r   Úschedule_collective_for_overlapz;_schedule_for_comm.<locals>.schedule_collective_for_overlapÂ   s‹   ø€ ô # 5Ô)Ð)Ð)Ù�Œà'¨Ñ.ˆà˜aÒÙ7Ó9Ð9�ÐFà�L‰L˜Ô#Ù�Y—_‘_Ô%Ø˜}¨Y¯_©_Ñ=Ñ=ˆOð ˜aÒÙ7Ó9Ð9�ÑFô
 	�‰�eÕr    z;Detected unscheduled nodes. Nodes with unmet dependencies: )Ú	enumeraterG   r,   r-   ÚsysÚmaxsizer   Ú	ancestorsrY   r   r   Úunmet_dependenciesr   Úestimate_op_runtimeÚitemsrI   rJ   rK   ÚaddÚheappopr)   )r   r   r   r   Úbuf_name_to_snodeÚidxr)   rL   Úop_nameÚ	node_nameÚcomm_idxÚancÚanc_fused_nameÚdepsrC   ra   r@   rM   r[   r2   rN   rQ   rO   r3   r4   r5   r`   rP   s                   @@@@@@@@@@@@r   r   r   I   sÈ  ÿû€ ðL ÐØÐØ#% r¨2Ð €Hˆh˜Ü Ó'ò "‰
ˆˆUØ×.Ñ.Ó0ò 	0ˆHØ*/Ð˜hÒ'ð	0ð ×0Ñ0Ó2ò 	0ˆGØ*/Ð˜wÒ'ð	0à/4Ð˜5Ÿ>™>Ó+Ñ,à—N‘NÓ$ˆ	Ü!Ÿk™kˆ�ÑØˆ�ÑØ!ˆ�Òð"ð €HØò +ˆÙÔ.¨uÔ5Ø)1ˆH�U—^‘^Ó%Ñ&Ø—‘ò S�Ø!3°CÑ!8×!AÑ!AÓ!C�Ü+.¨x¸Ñ/GÈÓ+R�˜Ò(ðSð ˜‰M‰HÚœM¨%Õ0Ø)*ˆH�U—^‘^Ó%Ò&ð+÷,ö ,ð  ö<àð 	ŒzÑG¨e×.FÑ.FÔGÓGÑGò<€Jð
 €EÜ=HÌÓ=T€LØDJÖK¸5�UÔ/°Ó6Ñ6ÒK€Mà!×'Ñ'Ó)ò )‰ˆˆtÜˆt‹9˜Š>Ü�N‰N˜5¡(¨5£/Ô2Øò 	)ˆCØ˜Ñ×!Ñ! %Õ(ñ	)ð)ð €I÷	;ð 	;ô6÷ô& ˆeŒ*Ü—‘˜eÓ$×*Ñ*ˆÙÔ#6°uÔ#=Ù+¨EÕ2á�UŒOô ˆe�*ð "×'Ñ'Ó)ò 
‰ˆˆtÜ�4‹y˜A‹~ð 	
ØIÈ*ÈÐVó	
ˆ~ð
ð ÐùòQ<ùò Ls   Ä9$JÅ5Jc                óx  — t         j                  j                  «       s| S | D �cg c]  }t        |«      sŒ|‘Œ }}t	        dt        |«      «      D ]a  }t        t        ||   j                  «       «      «      }||dz
     j                  «       D ]!  }||   j                  t        ||¬«      «       Œ# Œc | S c c}w )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   ÚrangerI   r*   r+   rG   Úadd_fake_depr   )ÚnodesÚname_to_bufr2   ÚnÚ
comm_nodesÚiru   Úbufs           r   Údecide_global_ordering_of_commsr�   ã   s¶   € ô ×Ñ×)Ñ)Ô+Øˆà"Ö=˜Ô&9¸!Õ&<’!Ð=€JÐ=ä�1”c˜*“oÓ&ò PˆäœD ¨A¡×!?Ñ!?Ó!AÓBÓCˆØ˜a !™eÑ$×5Ñ5Ó7ò 	PˆCØ�q‰M×&Ñ&¤w¨sÀÔ'NÕOñ	PðPð €Lùò >s
   ¥B7¶B7c                ó°   — t         j                  dk(  r| j                  «       }|S t        t         j                  «      sJ ‚t        j                  | «      }|S )z:
    Returns estimated op runtime in nanoseconds (ns)
    Údefault)r
   rg   Úget_estimated_runtimeÚcallable)r)   Úruntimes     r   rg   rg   ù   sR   € ô ×!Ñ! YÒ.Ø×-Ñ-Ó/ˆð €Nô œ×2Ñ2Ô3Ð3Ð3Ü×,Ñ,¨UÓ3ˆØ€Nr    c                ó¸  — d}t        | j                  t        j                  «      rd| j                  j                  › d�}d}| j                  j                  «       }t        |t        j                  «      rd|j                  › d|j                  › d�}| j                  j                  «       xs d}| j                  j                  j                  › |› |› d|› d�S )NÚ z (ú)z (size=z	, stride=)Ú
isinstanceÚnoder   ÚExternKernelOutÚpython_kernel_nameÚget_output_specÚLayoutÚsizeÚstrideÚmaybe_get_nameÚ	__class__r=   )r)   ÚdetailÚout_tensor_infoÚlayoutrn   s        r   Únode_summaryr—     sº   € Ø€FÜ�%—*‘*œb×0Ñ0Ô1Ø�e—j‘j×3Ñ3Ð4°AÐ6ˆØ€OØ�Z‰Z×'Ñ'Ó)€FÜ�&œ"Ÿ)™)Ô$Ø# F§K¡K =°	¸&¿-¹-¸ÈÐJˆØ—
‘
×)Ñ)Ó+Ò1¨r€IØ�j‰j×"Ñ"×+Ñ+Ð,¨V¨H°_Ð4EÀRÈ	À{ÐRSÐTÐTr    c                ó  — d}d }| D ]æ  }|€tt        |«      r|t        |«      z  }|j                  }n.t        |j                  «      rt	        d«      ‚|t        |«      z  }t
        j                  t        |«      › «       Œyt        |«      rt	        d«      ‚t        |j                  «      r"t
        j                  t        |«      › «       d }ŒÆt
        j                  dt        |«      › �«       Œè t
        j                  d|dz  dz  › �«       y )Ng        z8Wait is not expected when there is no collective runningzkFound two collectives running at the same time. `visualize_overlap` needs to be updated to handle this casez| zEst. runtime (ms): iè  )r   rg   r‹   r   ÚAssertionErrorÚoverlap_logÚdebugr—   )ÚorderÚtotal_est_runtimeÚcur_comm_noder)   s       r   Úvisualize_overlaprŸ     s  € Ø"ÐØ€MØò >ˆØÐ Ü" 5Ô)Ø!Ô%8¸Ó%?Ñ?Ð!Ø %§
¡
‘Ü˜Ÿ™Ô$Ü$ØNóð ð "Ô%8¸Ó%?Ñ?Ð!Ü×Ñ¤¨eÓ!4Ð 5Õ7ä" 5Ô)Ü$ðRóð ô ˜Ÿ™Ô$Ü×!Ñ!¤\°%Ó%8Ð$9Ô;Ø $‘ä×!Ñ! B¤|°EÓ':Ð&;Ð"<Õ=ð->ô. ×ÑØ
Ð/°$Ñ6¸Ñ=Ð>Ð?õr    c                ó‚  — | }t         j                  D ]À  }t        |t        «      r|t	        «       v rt	        «       |   }t
        j                  j                  «       dk(  r%t        j                  d|› d�«       	 t        |«        ||«      }t
        j                  j                  «       dk(  sŒœt        j                  d|› d�«       	 t        |«       ŒÂ |S # t        $ r(}t        j                  t        |«      «       Y d }~Œd }~ww xY w# t        $ r)}t        j                  t        |«      «       Y d }~�Œ&d }~ww xY w)Nr   z.==== Visualize overlap before reordering pass z ====z-==== Visualize overlap after reordering pass )r
   Ú'reorder_for_compute_comm_overlap_passesrŠ   ÚstrÚglobalsrv   rw   Úget_rankrš   r›   rŸ   Ú	Exception)r   rœ   ÚpÚes       r   Ú$reorder_compute_and_comm_for_overlapr¨   0  s  € ð €Eä×;Ñ;ò *ˆÜ�aœÔ !¤w£y¡.Ü“	˜!‘ˆAÜ×Ñ×%Ñ%Ó'¨1Ò,Ü×ÑØ@ÀÀÀ5ÐIôð*Ü! %Ô(ñ �%“ˆÜ×Ñ×%Ñ%Ó'¨1Ó,Ü×ÑØ?À¸sÀ%ÐHôð*Ü! %Õ(ð#*ð( €Løô ò *Ü×!Ñ!¤# a£&×)Ñ)ûð*ûô ò *Ü×!Ñ!¤# a£&×)Ò)ûð*ús0   Á:CÃ	DÃ	D	Ã!DÄD	Ä	D>ÄD9Ä9D>c           
     óÂ  ‡‡‡‡‡‡— t        | j                  «      Št        t         «      Št        t         «      Št        ‰«      D ]Ô  \  }}|j                  dk(  sŒ|j
                  t        j                  j                  j                  j                  k(  sŒR|j                  d   j                  dk(  sJ d|› d|j                  d   › d�«       ‚|j                  d   }|j                  d   }|dkD  r‰|   j                  |«       ŒÁ‰|   j                  |«       ŒÖ ˆˆˆfd„}t        t         «      }t        ‰«      D ]œ  \  }}|j                  dk(  sŒ|j
                  t        j                  j                  j                  j                  k(  sŒR|}|j                  d   Š‰j                  dk(  sJ d	‰› d
| › d�«       ‚ |‰«      sŒ‰|‰   j                  |«       Œž d„ }d„ Š‰D ]�  }|j                  dk(  sŒt        |j
                  t        j                   j"                  «      sŒB|j
                  j$                  j&                  sŒc ||«      rŒl ‰||j)                  «       «      sŒ„J d|› d�«       ‚ |j+                  «       D �]!  \  Š}	t        |	«      D �]  \  }
}‰|   }|j                  d   ‰u sJ ‚|j                  \  }Š|dz   }|
t-        |	«      dz
  k  r|	|
dz      nt-        ‰«      dz
  }‰|| }t/        ˆˆfd„|D «       «      rJ d‰› d|› d| › d�«       ‚|D ]ƒ  }|j                  dk(  sŒ‰|j                  v sŒ"|j
                  t        j                  j                  j                  j                  k7  sŒ^t1        ˆˆfd„|j                  D «       «      }||_        Œ… �Œ �Œ$ |j+                  «       D ].  \  Š}	t        |	«      D ]  \  }
}‰|   }| j3                  |«       Œ Œ0 ‰D ]q  }|j                  dk(  sŒ|j
                  t        j                  j                  j                  j                  k(  sŒO|j                  d   |v sŒa| j3                  |«       Œs y)aŠ  
    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	   c                ól  •— ‰j                  | g «      }‰j                  | g «      }t        |«      t        |«      k(  s2t        j                  d| › dt        |«      › dt        |«      › d�«       yt	        ||«      D ]7  \  }}||k\  sŒt        j                  d| › d‰|   › d|› d	‰|   › d|› d
�«        y y)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)ÚgetrI   Ú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_patternzLremove_fsdp2_unsharded_param_graph_input_usage.<locals>.check_resize_patternn  s  ø€ ð !J× MÑ MØ˜ó!
Ðð D×GÑGÈÐUWÓXÐäÐ(Ó)¬SÐ1CÓ-DÒDÜ�K‰KðHØHSÀ}ð UÜÐÓÐ ˜E¤#Ð&8Ó"9Ð!:ð ;ðôð ô 47Ø!Ð#5ó4
ò 	Ñ/Ð ð " _Ó4Ü—‘ðØ�Ð3°IÐ>PÑ4QÐ3RÐR\Ð]oÐ\pð q Ø )¨/Ñ :Ð;¸:ÀoÐEVð Wðôñ ð	ð r    z\
Assumed all FSDP2 `unsharded_param`s to be graph input, but it's not true!
Offending node: z	. Graph: c                óò   — | j                   t        j                  j                  j                  j
                  k(  xs; | j                   t        j                  j                  j                  j
                  k(  S r(   )Útargetrv   ÚopsÚfsdpÚcopy_rƒ   ÚinductorÚresize_storage_bytes_)r‹   s    r   Úis_allowed_mutationzKremove_fsdp2_unsharded_param_graph_input_usage.<locals>.is_allowed_mutationŸ  sO   € à�K‰Kœ5Ÿ9™9Ÿ>™>×/Ñ/×7Ñ7Ñ7ò OØ�{‰{œeŸi™i×0Ñ0×FÑF×NÑNÑNð	
r    c           	     ón  — t        | j                  t        j                  j                  «      r^t        | j                  j                  j                  «      D ��cg c])  \  }}|j                  �|j                  j                  r|‘Œ+ c}}ng }t        |D �cg c]5  }t        | j                  |   j                  d   j                  «       «      ‘Œ7 c}«      }t        |D �cg c](  }t        |j                  d   j                  «       «      ‘Œ* c}«      }t        ||z  «      dkD  S c c}}w c c}w c c}w )NÚvalr   )rŠ   r¼   rv   Ú_opsÚ
OpOverloadrb   Ú_schemaÚ	argumentsÚ
alias_infoÚis_writer   r   ÚargsÚmetaÚuntyped_storagerI   )r‹   Úunsharded_paramsr   rU   Úmutated_arg_idxesÚmutated_node_arg_storagesÚunsharded_paramÚstorages_of_unsharded_paramss           r   Ú-is_node_mutating_unsharded_param_or_its_aliaszeremove_fsdp2_unsharded_param_graph_input_usage.<locals>.is_node_mutating_unsharded_param_or_its_alias¥  s  € ô ˜$Ÿ+™+¤u§z¡z×'<Ñ'<Ô=ô & d§k¡k×&9Ñ&9×&CÑ&CÓD÷á�A�qØ—<‘<Ð+°·±×0EÒ0Eò ôð ð 	ô %/ð +öàô ˜tŸy™y¨™|×0Ñ0°Ñ7×GÑGÓIÕJòó%
Ð!ô (2ð (8öà#ô ˜×3Ñ3°EÑ:×JÑJÓLÕMòó(
Ð$ô Ð,Ð/KÑKÓLÈqÑPÐPùó)ùòùòs   Á.D'Â:D-Ã"-D2zdUser mutation on FSDP2 unsharded param is not allowed when Traceable FSDP2 is used. Violating node: c              3  ó2   •K  — | ]  } ‰|‰g«      –— Œ y ­wr(   © )rB   r‹   rÓ   rÑ   s     €€r   rD   zAremove_fsdp2_unsharded_param_graph_input_usage.<locals>.<genexpr>ê  s#   øè ø€ ò àñ >¸dÀ_ÐDU×Vñùs   ƒz(Assumed no ops mutating unsharded param z in subgraph z, but it's not true!
Graph: c              3  ó.   •K  — | ]  }|‰u r‰n|–— Œ y ­wr(   rÕ   )rB   ÚargÚreplacementrÑ   s     €€r   rD   zAremove_fsdp2_unsharded_param_graph_input_usage.<locals>.<genexpr>÷  s%   øè ø€ ò %àð (+¨oÑ'=™À3ÓFñ%ùs   ƒN)Úlistr{   r   rb   Úopr¼   rv   r½   rÀ   rÁ   rƒ   rË   rF   r¾   r¿   rŠ   rÅ   rÆ   rÇ   Ú
is_mutableÚkeysrh   rI   ÚanyÚtupleÚ
erase_node)Úgraphrl   r‹   r²   Únew_sizerº   Ú'unsharded_param_to_fsdp_copy_node_idxesÚfsdp_copy_noderÂ   Úfsdp_copy_node_idxesr   Úfsdp_copy_node_idxÚ_Úsubgraph_start_idxÚsubgraph_end_idxÚsubgraph_nodesÚnew_argsr·   r¸   rÓ   r¹   rØ   rÑ   s                    @@@@@@r   Ú.remove_fsdp2_unsharded_param_graph_input_usagerë   L  s�  ý€ ô �U—[‘[Ó!€Iô 1<¼DÓ0AÐ-Ü-8¼Ó->Ð*Ü˜yÓ)ò P‰	ˆˆTà�G‰G�Ó&Ø—‘œuŸy™y×1Ñ1×GÑG×OÑOÓOà—9‘9˜Q‘<—?‘? mÒ3ð ð :2Ø26°Ð7ZÐ[_×[dÑ[dÐefÑ[gÐZhð ið6ó Ð3ð Ÿ)™) A™,ˆKØ—y‘y ‘|ˆHØ˜!Š|Ø9¸+ÑF×MÑMÈcÕRà6°{ÑC×JÑJÈ3ÕOðPö"ôJ /:¼$Ó.?Ð+Ü˜yÓ)ò 	U‰	ˆˆTØ�7‰7�oÓ%¨$¯+©+¼¿¹¿¹×9MÑ9M×9UÑ9UÓ*UØ!ˆNØ"Ÿi™i¨™lˆOØ"×%Ñ%¨Ò6ð ð =à Ð! ¨5¨'ð 2ð9ó Ð6ñ $ OÕ4Ø7¸ÑH×OÑOÐPSÕTð	Uò
òQð4 ò ˆà�G‰G�Ó&Ü˜4Ÿ;™;¬¯
©
×(=Ñ(=Õ>Ø—‘×#Ñ#×.Ó.Ù'¨Õ-áDØÐ=×BÑBÓDõð ðeØeiÐdjð kðóð ðð: 
1×	6Ñ	6Ó	8ó")ñ 	ØØä%.Ð/CÓ%Dó 	)Ñ!ˆAÐ!Ø&Ð'9Ñ:ˆNØ!×&Ñ& qÑ)¨_Ñ<Ð<Ð<Ø+×0Ñ0‰NˆAˆ{à!3°aÑ!7Ðð ”sÐ/Ó0°1Ñ4Ò4ð % Q¨¡UÒ+ä˜“^ aÑ'ð ð
 'Ð'9Ð:JÐKˆNÜô à*ôô ð ð)Ø)8Ð(9¸À~ÐFVð WØ€wð ðóð ð 'ò 
)�à—G‘G˜Ó.Ø'¨4¯9©9Ò4ØŸ™¤u§y¡y×'9Ñ'9×'OÑ'O×'WÑ'WÓWä$ô %à#'§9¡9ô%ó  �Hð !)�D•Iò
)ò)	)ð	")ðP 
1×	6Ñ	6Ó	8ò-ñ 	ØØä%.Ð/CÓ%Dò 	-Ñ!ˆAÐ!Ø&Ð'9Ñ:ˆNØ×Ñ˜^Õ,ñ	-ð	-ð ò #ˆà�G‰G�Ó&Ø—‘œuŸy™y×1Ñ1×GÑG×OÑOÓOØ—	‘	˜!‘Ð GÒGà×Ñ˜TÕ"ñ#r    c                ó  ‡	— 	 dd l Š	‰	j                  j                  «       sJ ‚‰	j                  j                  j
                  r ‰	j                  j                  j                  sJ ‚	 ddl
m}m}m}m}m} 	 ˆ	fd„} |«       } | |‰	j                  j                  j
                  j                    |t"        j$                   |‰	j                  j&                  j(                  j                    |d«       |d«       |d«       |d«       |d	«       |d
«       |d«      «       |d«      «       |d«       |d«      «      |d„ ¬«      dˆ	fd„«       } || «       |j+                  | «       y # t        t        t        f$ r Y y w xY w)Nr   r	   )ÚCallFunctionÚ
KeywordArgÚMatchÚPatternMatcherPassÚregister_graph_patternc                óJ  •— t        | j                  «      }|D ]ˆ  }|j                  t        j                  k(  sŒ!|j
                  d   j                  ‰j                  j                  j                  j                  u sŒe|j
                  d   dk(  sŒx| j                  |«       ŒŠ y )Nr   r	   )rÙ   r{   r¼   ÚoperatorÚgetitemrË   r½   r¾   Úall_gather_copy_inrƒ   rß   )Úgr¹   r}   rv   s      €r   Úremove_unused_getitemz8reinplace_fsdp_all_gather.<locals>.remove_unused_getitem5  su   ø€ ä˜Ÿ™“Mˆ	Øò 	 ˆAà—‘œH×,Ñ,Ó,Ø—F‘F˜1‘I×$Ñ$¨¯	©	¯©×(IÑ(I×(QÑ(QÒQØ—F‘F˜1‘I “Nà—‘˜Q•ñ	 r    Úall_gather_inputsÚinp_split_sizesÚall_gather_input_numelÚ
world_sizeÚrankÚdtypeÚdeviceÚitem_idxÚ
group_sizeÚ
group_namec                ó&   — | j                   d   dk(  S )Nrÿ   r   )Úkwargs)Úmatchs    r   rV   z+reinplace_fsdp_all_gather.<locals>.<lambda>W  s   €  %§,¡,¨zÑ":¸aÑ"?€ r    )Ú	pass_dictÚextra_checkc                ó|   •— ˆfd„}| j                  ||d   |d   |d   |d   |d   |d   |d   |d	   |d
   g	«       y )Nc                 óú   •— | d d }| d   }| d   } ‰j                   j                  j                  j                  |Ž }|d   }|d   }‰j                   j                  j
                  j                  ||||¬«      }|S )Néþÿÿÿéÿÿÿÿr   r	   )Úout)r½   r¾   rõ   rƒ   Ú_c10d_functionalÚall_gather_into_tensor_out)	rË   Úcopy_in_argsr   r  rõ   rô   Ú	getitem_1Úall_gather_into_tensorrv   s	           €r   ÚreplzEreinplace_fsdp_all_gather.<locals>.reinplace_all_gather.<locals>.replZ  s–   ø€ ð    ˜9ˆLØ˜b™ˆJØ˜b™ˆJØ!J §¡§¡×!BÑ!B×!JÑ!JØð"Ðð )¨Ñ+ˆGØ*¨1Ñ-ˆIà—	‘	×*Ñ*×EÑE×MÑMØ˜Z¨¸ð Nó ð #ð
 *Ð)r    rø   rù   rú   rû   rü   rý   rþ   r   r  )Úreplace_by_example)r  rË   r  r  rv   s       €r   Úreinplace_all_gatherz7reinplace_fsdp_all_gather.<locals>.reinplace_all_gatherB  si   ø€ ô0	*ð$ 	× Ñ ØàÐ*Ñ+ØÐ(Ñ)ØÐ/Ñ0Ø�|Ñ$Ø�v‘Ø�w‘Ø�xÑ Ø�|Ñ$Ø�|Ñ$ð
õ	
r    )r  rï   )Ú5torch.distributed.fsdp._fully_shard._fsdp_collectivesrw   rx   r½   r  r  r  ÚImportErrorÚAttributeErrorr™   Úpattern_matcherrí   rî   rï   rð   rñ   rƒ   ró   rô   r¾   rõ   Úapply)
rà   rí   rî   rï   rð   rñ   r÷   Ú
graph_passr  rv   s
            @r   Úreinplace_fsdp_all_gatherr    ss  ø€ ð
ÛDà× Ñ ×-Ñ-Ô/Ð/Ð/ð �I‰I×&Ñ&×=Ò=Ø—	‘	×*Ñ*×EÒEð	
ðFØE÷
õ ðô 	 ñ $Ó%€JáÙØ�I‰I×&Ñ&×=Ñ=×EÑEÙÜ× Ñ ÙØ—I‘I—N‘N×5Ñ5×=Ñ=ÙÐ2Ó3ÙÐ0Ó1ÙÐ7Ó8Ù˜|Ó,Ù˜vÓ&Ù˜wÓ'Ù˜xÓ(ó	ñ ˜:Ó&óñ �|Ó$Ù�|Ó$ó#	
ð& Ù?ô+ô. 
ó/ð. 
ñD ˜%Ô Ø×Ñ�UÕøôE œ¬Ð8ò Ùðús   ƒA"E( Å(E?Å>E?c                óâ   — t        | t        j                  j                  j                  t        j                  j                  j
                  f«      rJ ‚t        | j                  «       dd  «      S )Né   )rŠ   rv   Ú	_inductorÚ	schedulerÚFusedSchedulerNodeÚGroupedSchedulerNodeÚintr-   )r)   s    r   Ú
get_op_idxr"    s]   € ÜØä�O‰O×%Ñ%×8Ñ8Ü�O‰O×%Ñ%×:Ñ:ð	
ôð ð ô ˆu�~‰~Ó  Ð#Ó$Ð$r    c           	     óz	  ‡‡‡ ‡!— ddl mŠ  g }t        t           «       }d}d}i }i }i Š!ˆ ˆ!fd„}	| D �]  }
t	        |
j
                  t        j                  j                  j                  j                  ¬«      �rßt        ˆfd„|
j                  D «       «      �rÀd}|
}t        «       }t        |||‰«       t        t        j                  j                  j                  j                  t        j                  j                  j                  j                  t        j                  j                  j                   j                  g«      Št#        |||‰ˆˆ fd„¬	«       t%        |d
„ ¬«      }t'        |«      }d}t)        t'        |«      «      D ]W  }||   }t+        |j
                  t        j                  j                  j                   j                  «      r|dz  }|dkD  sŒU|} n |d | }d }t)        t'        |«      dz
  «      D ]3  }t-        ||dz      j
                  t.        j0                  «      sŒ.|dz   } n |€J ‚ |	|d | «      } |	||d  «      }|||<   �Œ't+        |
j
                  t        j                  j                  j2                  j                  «      s�Œkd}|
}t        «       }t#        |||‰«       t%        |d„ ¬«      }d }t)        t'        |«      dz
  «      D ]3  }t-        ||dz      j
                  t.        j0                  «      sŒ.|dz   } n |€J ‚ |	|d | «      } |	||d  «      }|||<   �Œ t'        ‰!«      dkD  sJ ‚|rt'        |«      dkD  sJ ‚|rt'        |«      dkD  sJ ‚| D ]N  }
|
j5                  «       ‰!v r‰!|
j5                  «          }
|
|v rŒ-|j7                  |
«       |j9                  |
«       ŒP d }|j;                  «       D ]j  \  }}|�at=        t?        |jA                  «       «      «      }|jC                  «       D ],  }|jE                  tG        |j5                  «       |¬«      «       Œ. |}Œl d }|j;                  «       D ]j  \  }}|�at=        t?        |jA                  «       «      «      }|jC                  «       D ],  }|jE                  tG        |j5                  «       |¬«      «       Œ. |}Œl |S )Nr	   )r  Fc                ó˜   •— ‰j                   j                  | «      }| D ]  }|‰|j                  «       <   Œ |‰|j                  «       <   |S r(   )r   Úcreater-   )Úsnodes_to_groupÚ
group_noder)   r  Úsnode_name_to_final_snodes      €€r   Ú_create_group_nodez:enforce_comm_ordering_for_fsdp.<locals>._create_group_node™  sV   ø€ Ø×3Ñ3×:Ñ:¸?ÓKˆ
Ø$ò 	EˆEØ:DÐ% e§n¡nÓ&6Ò7ð	Eà;EÐ! *×"5Ñ"5Ó"7Ñ8ØÐr    )rÚ   c              3  ó¨   •K  — | ]I  }t        ‰|   j                  t        j                  j                  j
                  j                  «      –— ŒK y ­wr(   )r   r‹   rv   r½   r¾   rõ   rƒ   )rB   rU   r2   s     €r   rD   z1enforce_comm_ordering_for_fsdp.<locals>.<genexpr>¥  sD   øè ø€ ò 
ð ô Ø" 1Ñ%×*Ñ*¬E¯I©I¯N©N×,MÑ,M×,UÑ,U÷ñ
ùs   ƒAATc                ó–   •— t        | ‰j                  «      xs0 t        | ‰j                  «      xr | j                  j                  ‰v  S r(   )rŠ   ÚNopKernelSchedulerNodeÚExternKernelSchedulerNoder‹   Úop_overload)rU   Úallowed_opsr  s    €€r   rV   z0enforce_comm_ordering_for_fsdp.<locals>.<lambda>Ä  sF   ø€ Ü˜q )×"BÑ"BÓCò ä" 1 i×&IÑ&IÓJò >ØŸF™F×.Ñ.°+Ð=ð	'€ r    )Úcriteria_cbc                ó   — t        | «      S r(   ©r"  rT   s    r   rV   z0enforce_comm_ordering_for_fsdp.<locals>.<lambda>Ï  ó
   € ´J¸q³M€ r    rW   r   c                ó   — t        | «      S r(   r2  rT   s    r   rV   z0enforce_comm_ordering_for_fsdp.<locals>.<lambda>   r3  r    rt   )$rˆ   r  r   r   r   r‹   rv   r½   r  r  rƒ   rÝ   re   r   Úwait_tensorr¾   Úsplit_with_sizes_copyr   ÚsortedrI   ry   r   rŠ   r   Ú_WaitKernelÚ	chunk_catr-   rF   ri   rh   r*   r+   rG   Úget_outputsrz   r   )"r   r|   r2   Ú	new_orderrO   Ú	ag_existsÚ	rs_existsÚ$ag_grouped_node_to_wait_grouped_nodeÚ$rs_grouped_node_to_wait_grouped_noder)  r)   Ú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_noderu   ÚoÚprev_rs_waitr/  r  r(  s"     `                            @@@r   Úenforce_comm_ordering_for_fsdprR  Š  sñ  û€ õ
 à)+€IÜœ3‘Ó!€IØ€IØ€IØ+-Ð(Ø+-Ð(Ø "Ðõð ó nUˆäØ�J‰Jœ5Ÿ9™9×5Ñ5×PÑP×XÑXö
äó 
ð —_‘_ô	
õ 
ð ˆIØˆHÜLVËLÐ ô (ØØ$ØØ"ô	ô %ä—I‘I×.Ñ.×IÑI×QÑQÜ—I‘I×.Ñ.×:Ñ:×BÑBÜ—I‘I—N‘N×8Ñ8×@Ñ@ðóˆKô )ØØ$ØØ"ôõô !'Ø$Ñ*Aô!Ðô +.Ð.?Ó*@Ð'ØˆNÜœ3Ð0Ó1Ó2ò �Ø-¨aÑ0�	Ü!Ø—N‘N¤E§I¡I§N¡N×$HÑ$H×$PÑ$Pôð # aÑ'�NØ! AÓ%Ø23Ð/Ùðð !2Ð2NÐ3NÐ OÐð !ˆMÜœ3Ð0Ó1°AÑ5Ó6ò �ÜÐ/°°A±Ñ6×;Ñ;¼R¿^¹^ÕLØ$%¨¡E�MÙðð !Ð,Ð,Ð,Ù.Ð/@ÀÀ-Ð/PÓQˆMñ "4Ð4EÀmÀnÐ4UÓ!VÐàBTÐ0°Ó?ô ˜EŸJ™J¬¯	©	¯©×(@Ñ(@×(HÑ(HÖIØˆIØˆHô MWËLÐ Ü(ØØ$ØØ"ô	ô !'Ø$Ñ*Aô!Ðð
 !ˆMÜœ3Ð0Ó1°AÑ5Ó6ò �ÜÐ/°°A±Ñ6×;Ñ;¼R¿^¹^ÕLØ$%¨¡E�MÙðð !Ð,Ð,Ð,Ù.Ð/@ÀÀ-Ð/PÓQˆMñ "4Ð4EÀmÀnÐ4UÓ!VÐàBTÐ0°Ó?ð]nUô` Ð(Ó)¨AÒ-Ð-Ð-ÙÜÐ7Ó8¸1Ò<Ð<Ð<ÙÜÐ7Ó8¸1Ò<Ð<Ð<ð ò ˆØ�>‰>ÓÐ8Ñ8Ø-¨e¯n©nÓ.>Ñ?ˆEØ�IÑØØ×Ñ˜ÔØ�‰�eÕðð €LØ*N×*TÑ*TÓ*Vò 'Ñ&ˆ�ØÐ#Ü¤ ]×%CÑ%CÓ%EÓ FÓGˆLØ!×-Ñ-Ó/ò �Ø×*Ñ*Ü˜AŸJ™J›L°|ÔDõðð '‰ð'ð €LØ*N×*TÑ*TÓ*Vò 'Ñ&ˆ�ØÐ#Ü¤ ]×%CÑ%CÓ%EÓ FÓGˆLØ!×-Ñ-Ó/ò �Ø×*Ñ*Ü˜AŸJ™J›L°|ÔDõðð '‰ð'ð Ðr    )r   úlist[BaseSchedulerNode]r;   rS  )
r   rS  r   Úboolr   rT  r   rT  r;   rS  )r{   rS  r;   rS  )r)   r   r;   Úfloat)rà   útorch.fx.Graph)rà   rV  r;   r<   )r   ú1list[torch._inductor.scheduler.BaseSchedulerNode]r|   z4dict[str, torch._inductor.scheduler.SchedulerBuffer]r2   zdict[str, BaseSchedulerNode]r;   rW  )2Ú
__future__r   rJ   Úloggingró   rc   Úcollectionsr   Útypingr   r   rv   Ú torch.multiprocessing.reductionsr   Útorch.utils._ordered_setr   rˆ   r
   r   Údependenciesr   Úutilsr   r   r   r   r   r   r   Ú	getLoggerr=   r¯   Ú_loggingÚgetArtifactLoggerrš   r  r   r   r   r#   r   r�   rg   r—   rŸ   r¨   rë   r  r"  rR  rÕ   r    r   ú<module>rc     s>  ðõ #ã Û Û Û 
Ý #ß %ã Ý ;Ý /ç Ý !÷÷ ñ ð €g×Ñ˜Ó!€Ø�n‰n×.Ñ.¨x¸ÓC€áÝ,óóðØ#ðàóð&WØ#ðWàðWð ðWð ð	Wð
 óWðtØ"ðàóó,	ò	Uòð>Ø#ðàóó8A#óHlò^%ðnØ=ðnàEðnð 5ðnð 7ô	nr    