ó
    Eñi0  ã                  ó  • % S SK Jr  S SKrS SKrS SKJrJrJr  S SKJ	r	  S SK
r
S SKrS SKJr  S SKrS SKJr  SSK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J r   SS
K!J"r"  SSK#J$r$J%r%  \(       a  S SK&J'r'J(r(  Sr)Sq*S\+S'   \RX                   " S S5      5       r-\RX                   " S S5      5       r.SS jr/S S jr0\Rb                  S!S j5       r2          S"S jr3 " S S\5      r4S#S jr5 " S S5      r6    S$S jr7      S%S jr8g)&é    )ÚannotationsN)ÚAnyÚTYPE_CHECKINGÚUnion)Úpatch)Ú
OrderedSeté   )ÚconfigÚselect_algorithm)ÚBufferÚChoiceCallerÚLayoutÚMultiTemplateBufferÚOperationBufferÚ
StorageBoxÚ	TensorBox)ÚKernelInputsÚMMKernelInputs)ÚSchedulerNode)ÚNullHandlerÚV)Ú	GeneratorÚSequenceÚdistributed_autotuneúdist.ProcessGroup | NoneÚ_AUTOTUNE_PGc                  ó6   • \ rS rSr% SrSrS\S'   SrS\S'   Srg)	Ú_DistributedAutotuneStateé'   z9
State used to track autotuning during a graph_context()
r   ÚintÚautotuned_indexÚautotuned_local_count© N)	Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r!   Ú__annotations__r"   Ú__static_attributes__r#   ó    Úa/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/_inductor/distributed_autotune.pyr   r   '   s    ‡ ñð €O�SÓð "#Ð˜3Ö"r+   r   c                  ó*   • \ rS rSr% S\S'   S\S'   Srg)Ú_DistributedAutotuneInfoé6   r    ÚindexÚboolÚlocalr#   N)r$   r%   r&   r'   r)   r*   r#   r+   r,   r.   r.   6   s   ‡ àƒJØ†Kr+   r.   c                 óÀ   • [         R                  " 5       (       aD  [         R                  " 5       (       a*  [        c  [         R                  R                  SS9q[        $ g )NÚpt2_distributed_autotune_pg)Úpg_tag)ÚdistÚis_availableÚis_initializedr   Údistributed_c10dÚ_new_group_with_tagr#   r+   r,   Úget_autotune_pgr;   <   sO   € Ü×Ò×Ñœt×2Ò2×4Ñ4äÑÜ×0Ñ0×DÑDØ4ð Eð ˆLô Ðàr+   c                ót   • [         R                  (       d   e[        U 5      n[        U5      n[	        X5        g)z”
Finish the distributed autotuning by propagating the autotuning results
between the ranks and then replacing the placeholder with the real Buffer.
N)r
   Údistributed_max_autotune_gemmÚ_autotune_local_nodesÚ_syncÚ_autotune_remote_nodes)Ú	schedulerÚautotune_resultsÚchoices_by_indexs      r,   ÚschedulerD   H   s2   € ô
 ×/×/Ð/Ð/Ü,¨YÓ7ÐÜÐ-Ó.ÐÜ˜9Õ7r+   c               #  ó&  #   • [        [        R                  " SS9[        5      (       a   e[        R                  " [        5       5         Sv •  [        R                  " [        5       5        g! [        R                  " [        5       5        f = f7f)zX
Wrapped around processing a graph, sets up figuring out which ranks tune
which shapes.
F)Úcheck_poisonedN)Ú
isinstancer   Úget_distributed_autotune_stater   Úset_distributed_autotune_stater   r#   r+   r,   Úgraph_contextrJ   S   sk   é € ô Ü	×(Ò(¸Ñ>Ü!÷ñ ð ð ô ×$Ò$Ô%>Ó%@ÔAð8Ûä	×(Ò(¬«Õ7øŒ×(Ò(¬«Õ7üs   ‚ABÁA. ÁBÁ. BÂBc                ó"  • [         R                  (       d  g[        5       =n(       d  g[        U5      S::  a  g[        R
                  nUR                  nU=R                  S-  sl        XdR                  5       -  UR                  5       :H  n[        Xg5      [        R                  R                  [        '   U(       a  U=R                  S-  sl        g[        R                  R                   R"                  R%                  ['        XU5      5      $ )z‰
Used by an op (like `mm`) to determine if the op should be autotuned
locally (returns None) or remotely (returns a placeholder Buffer).
Nr	   )r
   r=   r;   Úlenr   Údistributed_autotune_stater!   ÚsizeÚrankr.   Úcurrent_nodeÚmetaÚ_DISTRIBUTED_AUTOTUNE_KEYr"   ÚtorchÚ	_inductorÚirr   ÚcreateÚ_DistributedAutotuneBuffer)ÚnameÚchoicesÚinputsÚlayoutÚautotune_pgÚstater0   r2   s           r,   Úmaybe_autotune_remoter^   d   sÙ   € ô ×/×/Øä*Ó,Ð,ˆKÕ,Øä
ˆ7ƒ|�qÓØä×(Ñ(€EØ×!Ñ!€EØ	×Ò˜QÑÕØ×$Ñ$Ó&Ñ&¨+×*:Ñ*:Ó*<Ñ<€Eä5MØó6„A‡N�N×ÑÔ1Ñ2ö Ø×#Ò# qÑ(Õ#Øä�?‰?×Ñ×'Ñ'×.Ñ.Ü" 4°Ó8óð r+   c                  óh   ^ • \ rS rSr% SrS\S'           S	U 4S jjr    S
S jrSS jrSr	U =r
$ )rW   é…   zœ
A MultiTemplateBuffer which represents a kernel being autotuned on a
different rank. When `schedule` is called this will be replaced by the
"real" buffer.
ÚstrÚ_kernel_namec           	     óZ   >• [         TU ]  UUU R                  / [        0 5      S9  Xl        g )N)Úchoice_timings_fnÚunfiltered_choicesÚallowed_prologue_inps)ÚsuperÚ__init__Ú_dummy_choice_timingsr   rb   )ÚselfÚkernel_namerZ   r[   Ú	__class__s       €r,   rh   Ú#_DistributedAutotuneBuffer.__init__�   s8   ø€ ô 	‰ÑØØØ"×8Ñ8Ø!Ü",¨R£.ð 	ñ 	
ð (Õr+   c                ó   • [         e©N)ÚNotImplementedError)rj   Ú_hint_overrides     r,   ri   Ú0_DistributedAutotuneBuffer._dummy_choice_timingsŸ   s
   € ô
 "Ð!r+   c                óÆ  • SSK Jn  [        R                  " [        R
                  SS5         [        / U R                  Q5      n[        U R                  [        5      (       d   eUR                  U R                  U5      nU" U R                  U/UR                  5       U R                  5      n[        U[        5      (       d   eUsSSS5        $ ! , (       d  f       g= f)z]
Given a _SerializedChoice (autotune results from another rank)
compute the final TensorBox.
r	   )Úautotune_select_algorithmrA   N)r   rt   r   Úobjectr   Úgraphr   Úoriginal_inputsrG   r[   r   Ú
get_choicerb   Únodesr   )rj   Ú
ser_choicert   Úkernel_inputsÚchoiceÚbuffers         r,   ÚautotuneÚ#_DistributedAutotuneBuffer.autotune¦   s®   € õ 	@ä�\Š\œ!Ÿ'™' ;°Õ5Ü*Ð+B¨T×-AÑ-AÐ+BÓCˆMÜ˜dŸk™k¬6×2Ñ2Ð2Ð2Ø×*Ñ*¨4¯;©;¸ÓFˆFÙ.Ø×!Ñ!Ø�Ø×#Ñ#Ó%Ø—‘ó	ˆFô ˜f¤i×0Ñ0Ð0Ð0Ø÷ 6×5×5ús   ­BCÃ
C )rb   )rk   ra   rZ   úlist[Buffer]r[   r   ÚreturnÚNone)rq   z
int | Noner�   zdict[ChoiceCaller, float])rz   Ú_SerializedChoicer�   r   )r$   r%   r&   r'   r(   r)   rh   ri   r~   r*   Ú__classcell__)rl   s   @r,   rW   rW   …   sZ   ø‡ ñð Óð(àð(ð ð(ð ð	(ð
 
÷(ð "Ø(ð"à	"ô"÷ò r+   rW   c                ó‚  • [        5       nU(       d   eS/UR                  5       -  n[        R                  R	                  X US9  [        S U 5       5      nS/U-  nSnU HG  nU H>  n[        U[        5      (       d   eXGR                     b   eXtUR                  '   US-  nM@     MI     X5:X  d   SU SU 35       eU$ )zL
Perform the all_gather to collect the autotune results from all the ranks.
N)Úgroupc              3  ó8   #   • U  H  n[        U5      v •  M     g 7fro   )rL   )Ú.0Úxs     r,   Ú	<genexpr>Ú_sync.<locals>.<genexpr>É   s   é € Ð0¢Z ”S˜—V�V¢Zùs   ‚r   r	   zcount mismatch: ú != )	r;   rN   rS   ÚdistributedÚall_gather_objectÚsumrG   rƒ   r0   )rB   r\   Ú
all_statesÚ
node_countrC   Úcheck_countÚother_resultsr|   s           r,   r?   r?   ½   sÞ   € ô
 "Ó#€KÞÐˆ;ð 26°¸×9IÑ9IÓ9KÑ0K€JÜ	×Ñ×'Ñ'¨
ÈKÐ'ÑXäÑ0¡ZÓ0Ó0€Jà15°¸Ñ0CÐà€KÛ#ˆÛ#ˆFÜ˜fÔ&7×8Ñ8Ð8Ð8Ø#§L¡LÑ1Ñ9Ð9Ð9Ø-3˜VŸ\™\Ñ*Ø˜1ÑŠKó	 $ñ $ð Ó$ÐVÐ(8¸¸ÀDÈÈÐ&VÓVÐ$ØÐr+   c                  ó^   • \ rS rSrSrS
S jrSS jr\SS j5       r\SS j5       r	SS jr
Srg	)rƒ   éÙ   zÆ
This is a serializer for the autotune choice. KernelTemplateChoice can't
be serialized directly (the template and inputs prevent this) so we need to
serialize it by parts and reconstruct later on.
c                ó„   • Xl         [        R                  U5      U l        U R	                  UR
                  5      U l        g ro   )r0   rƒ   Ú_template_uid_from_choiceÚtemplate_uidÚ_compute_kwargsÚdescriptionÚkwargs)rj   r0   r|   s      r,   rh   Ú_SerializedChoice.__init__à   s2   € ØŒ
Ü-×GÑGÈÓOˆÔØ×*Ñ*¨6×+=Ñ+=Ó>ˆ�r+   c                ó&  • U R                  5       n0 U R                  EnSU;   aF  UR                  5       S   R                  5       S   n[        R
                  " XTS   5      US   :H  US'   0 nSSKJnJn  U" U5      n	U" X9XaU5      n
U
R                  $ )z-
Deserialize the ChoiceCaller and return it.
ÚBLOCK_Kr   r	   ÚEVEN_K)ÚDictKernelTemplateParamsÚKernelTemplateChoice)
Ú_template_from_uidr›   ry   Úget_sizeÚsympyÚgcdÚkernel_template_choicer    r¡   r|   )rj   r[   rZ   Útemplater›   ÚkÚextra_kwargsr    r¡   ÚparamsÚktcs              r,   rx   Ú_SerializedChoice.get_choiceå   s—   € ð
 ×*Ñ*Ó,ˆà �D—K‘K�ˆØ˜Óð
 —‘“˜qÑ!×*Ñ*Ó,¨QÑ/ˆAÜ$Ÿyšy¨°9Ñ,=Ó>À&ÈÑBSÑSˆF�8Ñà')ˆ÷	
ñ
 *¨&Ó1ˆÙ" 8°\È6ÓRˆØ�z‰zÐr+   c                ó”  • U (       d  0 $ 0 nU R                  S5       H§  nUR                  SS5      u  p4UR                  5       UR                  5       pCUS:X  a  SX'   MB  US:X  a  SX'   MN  UR                  5       (       a  [        U5      X'   Mr  UR	                  S5      (       a  UR                  S5      (       d   eUSS	 X'   M©     U$ )
z9
Given a template description turn it into input kwargs.
Ú,Ú=r	   ÚTrueTÚFalseFÚ'éÿÿÿÿ)ÚsplitÚstripÚisdigitr    Ú
startswithÚendswith)rš   r›   ÚcfgÚkeyÚvals        r,   r™   Ú!_SerializedChoice._compute_kwargsÿ   s¶   € ö
 ØˆIð 46ˆØ×$Ñ$ SÖ)ˆCØ—y‘y  aÓ(‰HˆCØ—y‘y“{ C§I¡I£K�Ø�f‹}Ø"�“Ø˜“Ø#�“Ø—‘—‘Ü! #›h�“à—~‘~ c×*Ñ*¨s¯|©|¸C×/@Ñ/@Ð@Ð@Ø! ! B˜i�“ñ *ð ˆr+   c                ó*  • [        U [        R                  5      (       a>  U R                  R                  S:X  a  g[        SU R                  R                  < 35      e[        U [        R                  5      (       a  g[        S[        U 5       35      e)zi
Given a ChoiceCaller figure out which template represents it. This
is reversed by _template_from_uid().
Úmmz!torch._inductor.kernel.mm.aten_mmzTODO: kernel z%torch._inductor.kernel.mm.mm_templatezTODO: )rG   r   ÚExternKernelCallerr|   rX   ÚRuntimeErrorÚTritonTemplateCallerÚtype)r|   s    r,   r—   Ú+_SerializedChoice._template_uid_from_choice  sw   € ô �fÔ.×AÑA×BÑBØ�}‰}×!Ñ! TÓ)Ø:ä" ]°6·=±=×3EÑ3EÑ2HÐ#IÓJÐJÜ˜Ô 0× EÑ E×FÑFØ:ä ¬¨V« ~Ð6Ó7Ð7r+   c                óŠ   • U R                   R                  S5      n[        5       US      nUSS  H  n[        X#5      nM     U$ )z"
See _template_uid_from_choice().
Ú.r   r	   N)r˜   r´   ÚglobalsÚgetattr)rj   ÚpartsÚobjr¨   s       r,   r¢   Ú$_SerializedChoice._template_from_uid+  sH   € ð ×!Ñ!×'Ñ'¨Ó,ˆÜ‹i˜˜a™Ñ!ˆØ�q�r“ˆAÜ˜#“/ŠCñ àˆ
r+   )r0   r›   r˜   N)r0   r    r|   r   r�   r‚   )r[   r   rZ   r   r�   zChoiceCaller | None)rš   ra   r�   z dict[str, Union[int, str, bool]])r|   r   r�   ra   )r�   r   )r$   r%   r&   r'   r(   rh   rx   Ústaticmethodr™   r—   r¢   r*   r#   r+   r,   rƒ   rƒ   Ù   s>   † ñô?ô
ð4 óó ðð0 ó8ó ð8÷$r+   rƒ   c                ó€  • / nU R                    Há  n[        U[        5      (       d  M  UR                  =nc  M+  [        U[        5      (       a  MB  [        U[
        5      (       d  MY  UR                  =nc  Mj  UR                  =nc  M{  UR                  [        5      nUc  M•  UR                  (       d   eUR                  5       u  px[        UR                  U5      n	UR                  U	5        Mã     [        R                   n
[#        U5      U
R$                  :X  d!   S[#        U5       SU
R$                   S35       eU$ )zh
Go through the nodes in the scheduler and autotune the kernels which
should be autotuned by this rank.
z'incorrect local autotuned nodes found (rŒ   Ú))ry   rG   r   ÚnoderW   r   Úorigin_noderQ   ÚgetrR   r2   Úget_min_choicerƒ   r0   Úappendr   rM   rL   r"   )rA   rB   rÎ   Ú
inner_noderÏ   rQ   ÚinfoÚ
min_choiceÚ_r|   r]   s              r,   r>   r>   6  s0  € ð 13Ðà—”ˆÜ˜$¤×.Ñ.ÙàŸ)™)Ð#ˆJÑ,Ùä�jÔ"<×=Ñ=áä˜*Ô&9×:Ñ:Ùà%×1Ñ1Ð1ˆKÑ:Ùà×$Ñ$Ð$ˆDÑ-Ùà�x‰xÔ1Ó2ˆØ‰<Ùà�z�zÐˆzð
 #×1Ñ1Ó3‰ˆ
ä" 4§:¡:¨zÓ:ˆØ×Ñ Ö'ñA  ôD ×(Ñ(€EÜÐÓ  E×$?Ñ$?Ó?ð Ø
1´#Ð6FÓ2GÐ1HÈÈU×MhÑMhÐLiÐijÐkóÐ?ð Ðr+   c                ó.  • [        U R                  5       Hü  u  p#[        U[        5      (       d  M  [        UR                  =n[
        5      (       d  M?  UR                  c   eUR                  R                  [           nUR                  XR                     5      nUR                  n[        U[        5      (       d   eUR                  n[        U[        5      (       d   eUR                  UR                  :X  d   eU R                  X„X#5        Mþ     g)zc
Go through the nodes in the scheduler and autotune the nodes that were
autotuned on remote ranks.
N)Ú	enumeratery   rG   r   rÎ   rW   rÏ   rQ   rR   r~   r0   Údatar   r   r[   Ú_replace_node)	rA   rC   ÚirÎ   Ú	dist_noderÔ   Úout_tensorboxÚout_storageÚ
out_buffers	            r,   r@   r@   i  sê   € ô ˜YŸ_™_Ö-‰ˆÜ�dœM×*Ó*¬zØŸ)™)Ð#ˆYÔ&@÷0
ó 0
ð ×(Ñ(Ñ4Ð4Ð4Ø×(Ñ(×-Ñ-Ô.GÑHˆDØ%×.Ñ.Ð/?Ç
Á
Ñ/KÓLˆMà'×,Ñ,ˆKÜ˜k¬:×6Ñ6Ð6Ð6Ø$×)Ñ)ˆJÜ˜j¬/×:Ñ:Ð:Ð:à×$Ñ$¨	×(8Ñ(8Ó8Ð8Ð8à×#Ñ# J¸1ÖCò .r+   )r�   r   )rA   ú#torch._inductor.scheduler.Schedulerr�   r‚   )r�   zGenerator[None, None, None])
rX   ra   rY   zlist[ChoiceCaller]rZ   r€   r[   r   r�   zTensorBox | None)rB   úlist[_SerializedChoice]r�   úSequence[_SerializedChoice])rA   rà   r�   rá   )rA   rà   rC   râ   r�   r‚   )9Ú
__future__r   Ú
contextlibÚdataclassesÚtypingr   r   r   Úunittest.mockr   r¤   Útorch._loggingrS   Útorch.distributedr�   r6   Útorch.fxÚtorch.utils._ordered_setr   Ú r
   r   rU   r   r   r   r   r   r   r   r{   r   r   rA   r   Úvirtualizedr   r   Úcollections.abcr   r   rR   r   r)   Ú	dataclassr   r.   r;   rD   ÚcontextmanagerrJ   r^   rW   r?   rƒ   r>   r@   r#   r+   r,   Ú<module>rñ      sG  ðÞ "ã Û ß ,Ñ ,Ý ã ã Ý  Û Ý /ç &÷÷ ñ ÷ 8Ý $ß 'ö ß3ð 3Ð à)-€Ð&Ó -ð ×Ñ÷#ð #ó ð#ð ×Ñ÷ð ó ðô
	ô8ð ×Ñó8ó ð8ð Ø
ðØ*ðØ4@ðØJPðàôôB4Ð!4ô 4ôp÷8Zñ Zðz0Ø2ð0àô0ðfDØ2ðDà1ðDð 
õDr+   