Ë
    *êñiæ.  ã                  ó¨  — U d Z ddlmZ ddlZddlZddlmZ ddlmZ ddl	m
Z
mZmZmZmZ erddlmZmZ ddlZddlmZ g d¢Zd	ed
<    ej0                  e«      Z ed«      Ze G d„ dee   «      «       Zdddddœ	 	 	 	 	 	 	 	 	 	 	 dd„Z	 	 d	 	 	 	 	 	 	 dd„Zd„ f	 	 	 	 	 	 	 	 	 dd„Zdd„Z 	 	 	 	 	 	 dd„Z!	 	 	 	 	 	 dd„Z"	 	 	 	 	 	 dd„Z#d d„Z$d!d„Z%y)"zl
A set of primitive functions for performing collective ops.

Each should also handle single rank scenario.
é    )ÚannotationsN)Údefaultdict)Ú	dataclass)ÚAnyÚcastÚGenericÚTYPE_CHECKINGÚTypeVar)ÚCallableÚIterable)ÚSyncPayloadÚ	broadcastÚ
all_gatherÚall_gather_object_enforce_typez	list[str]Ú__all__ÚTc                  ó:   — e Zd ZU ded<   ded<   ded<   dZded	<   y)
r   ú
str | NoneÚ
stage_nameÚboolÚsuccessr   ÚpayloadNzException | NoneÚ	exception)Ú__name__Ú
__module__Ú__qualname__Ú__annotations__r   © ó    úd/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/torch/distributed/collective_utils.pyr   r   &   s   … àÓØƒMØƒJØ"&€IÐÔ&r   r   T)r   r   ÚrankÚpgc               óT  — |s| �t        d«      ‚d}d}|€|dk(  s|�)|j                  «       |k(  rt        | «      r		  | «       }n| }t	        ||||¬«      }|�E|g}	t        j                  |	||¬«       t        |	«      dk7  rt        dt        |	«      › �«      ‚|	d   }|j                  sNd	|› d
�}
|�|
d|j                  › �z  }
|j                  �|
d|j                  › �z  }
t        |
«      |j                  ‚t        t        |j                  «      S # t        $ r}d}|}Y d}~ŒÜd}~ww xY w)aK  
    Broadcasts the data payload from rank 0 to all other ranks.
    Or if a function is passed, execute it in rank 0 and broadcast result to all other ranks.

    Can be used to broadcast a failure signal to stop all ranks.

    If the function raises an exception, all ranks will raise.

    Args:
        data_or_fn: the data to broadcast or function to execute and broadcast result.
        success: False to stop all ranks.
        stage_name: the name of the logical stage for synchronization and debugging
        rank: rank to broadcast data or execute function and broadcast results.
        pg: the process group for sync
    Throws:
        RuntimeError from original exception trace
    Returns:
        the value after synchronization

    Example usage:
    >> id = broadcast(data_or_fn=allocate_id, rank=0, pg=ext_pg.my_pg)
    Nz9Data or Function is expected to be None if not successfulr   F©r   r   r   r   )ÚsrcÚgroupé   z7Expected broadcast_list to have exactly 1 element, got zRank z failedz: stage z: exception )ÚAssertionErrorr!   ÚcallableÚ	Exceptionr   ÚdistÚbroadcast_object_listÚlenr   r   r   ÚRuntimeErrorr   r   r   )Ú
data_or_fnr   r   r!   r"   r   r   ÚeÚsync_objÚbroadcast_listÚ	error_msgs              r    r   r   .   sk  € ñ> �zÐ-ÜØGó
ð 	
ð €GØ"&€Ià
ˆ
�t˜q’y b n¸¿¹»ÀdÒ9Jä�JÔðÙ$›,‘ð
 !ˆGô ØØØØô	€Hð 
€~Ø"˜ˆÜ×"Ñ" >°tÀ2ÕFÜˆ~Ó !Ò#Ü ØIÌ#ÈnÓJ]ÐI^Ð_óð ð " !Ñ$ˆð ×ÒØ˜D˜6 Ð)ˆ	ØÐ!Ø˜8 H×$7Ñ$7Ð#8Ð9Ñ9ˆIØ×ÑÐ)Ø˜<¨×(:Ñ(:Ð';Ð<Ñ<ˆIä˜9Ó%¨8×+=Ñ+=Ð=ä”�8×#Ñ#Ó$Ð$øôC ò Ø�Ø•	ûðús   ¼D Ä	D'ÄD"Ä"D'c                ó8  — d}d}d}t        | «      r		  | «       }n| }t        ||||¬«      }|��dgt        j                  |«      z  }t        |||«       t        t        t           |d   «      j                  }g }	g }
d}t        t        t        t        t              |«      «      D ]|  \  }}|j                  |k7  r|d|› d|j                  › d	�z  }Œ,|j                  s*|j                  �|	j                  ||j                  f«       Œb|
j                  |j                  «       Œ~ t        |	«      dkD  rt!        ||	«      |	d   ‚|
S |j                  s#t!        d
|j                  › �«      |j                  ‚|j                  gS # t        $ r}d}|}Y d}~�Œwd}~ww xY w)a.  
    A simple all_gather primitive with basic synchronization guard logic,
    by checking payload from all ranks has the same stage name.

    Args:
        data_or_fn: the data to be all gathered across ranks or function to be executed
        stage_name: the sync stage name for out-of-sync protection
        pg: the process group for sync
    Throws:
        RuntimeError from original exception trace
    Returns:
        a list of synced data from all ranks

    Example usage:
    >> all_ids = all_gather(data_or_fn=allocate_id, pg=ext_pg.my_pg)
    NTFr$   r   Ú z)Unexpected stage name received from rank z: ú z!all_gather failed with exception )r)   r*   r   r+   Úget_world_sizer   r   r   r   Ú	enumerateÚlistr   r   Úappendr   r-   r.   )r/   r   r"   r   r   r   r0   r1   Ú
total_listÚexception_listÚret_listr3   ÚiÚsps                 r    r   r   ~   s¼  € ð* €GØ"&€IØ€Gä�
Ôð	Ù “l‰Gð
 ˆäØØØØô	€Hð 
�~à�Vœd×1Ñ1°"Ó5Ñ5ˆ
Ü& r¨:°xÔ@äœ+¤a™.¨*°Q©-Ó8×CÑCˆ
Ø68ˆØˆØˆ	äœt¤D¬´Q©Ñ$8¸*ÓEÓFò 		(‰EˆAˆrØ�}‰} 
Ò*ØØ?À¸sÀ"ÀRÇ]Á]ÀOÐSTÐUñ�	ð Ø—:’: "§,¡,Ð":Ø×%Ñ% q¨"¯,©,Ð&7Ô8ØØ�O‰O˜BŸJ™JÕ'ð		(ô ˆ~Ó Ò"ÜØØóð " !Ñ$ð%ð ˆà×ÒÜØ3°H×4FÑ4FÐ3GÐHóà×%Ñ%ð&ð × Ñ Ð!Ð!øô[ ò 	ØˆGØŽIûð	ús   “F Æ	FÆ
FÆFc                ó.   — t        | «      t        |«      u S )N)Útype)ÚxÚys     r    ú<lambda>rD   Ô   s   € ¼DÀ»GÄtÈAÃwÐ<N€ r   c                óì   — t        j                  ||| ¬«       t        |«      }|dk(  ry|d   }t        d|«      D ]7  } ||||   «      rŒt	        d|› dt        ||   «      › dt        |«      › �«      ‚ y)aN  
    Similar to plain all_gather_object but with additional type checking
    AFTER gather is done to ensure basic consistency.
    If check does not pass, all ranks will fail with exception.

    This is generally to prevent conditional logic leading to
    unexpected messages being received. This is considered fatal code error,
    but due to logic stacks this might happen implicitly in practice.

    The default check does not check sub type (considered different)
    or covariance (considered same) but users can pass in custom checker
    if more complicated check is needed.
    )r&   r   Nr'   zObject type at index z is z, while first object type is )r+   Úall_gather_objectr-   ÚrangeÚ	TypeErrorrA   )r"   Úobject_listÚobjÚtype_checkerÚlist_lenÚ	first_objr>   s          r    r   r   Í   s“   € ô, 	×Ñ˜;¨°2Õ6ô �;Ó€HØ�1‚}ØØ˜A‘€IÜ�1�hÓò ˆÙ˜I {°1¡~Õ6ÜØ'¨ s¨$¬t°KÀ±NÓ/CÐ.Dð E.Ü.2°9«oÐ->ð@óð ñr   c                óV  — t        | «      } t        | «      dk  rt        d«      ‚t        t	        | «      «      t        | «      k7  rt        d«      ‚d }g }| rÎ| j                  d«      }|€|}nµt        |t        «      r/||dz   k(  rt        ||dz   d«      }nŒ||z
  }t        |||z   |«      }nvt        |t        «      st        d«      ‚||j                  k(  r9t        |j                  |j                  |j                  z   |j                  «      }n|j                  |«       |}| rŒÎt        |t        «      r |j                  t        ||dz   d«      «       n!t        |t        «      r|j                  |«       g }|D ]ž  }t        |«      dk(  r|j                  |j                  › «       Œ.|j                  dk(  r+|j                  |j                  › d|j                  › �«       Œh|j                  |j                  › d|j                  › d|j                  › �«       Œ  dj                  |«      S )Nr   zranks should all be positivez#ranks should not contain duplicatesr'   z!curr must be an instance of rangeú:ú,)ÚsortedÚminr(   r-   ÚsetÚpopÚ
isinstanceÚintrG   ÚstopÚstartÚstepr:   Újoin)ÚranksÚcurrÚrangesrB   rY   ÚresultÚrs          r    Ú_summarize_ranksr`   ò   sÔ  € Ü�5‹M€EÜ
ˆ5ƒz�A‚~ÜÐ;Ó<Ð<Ü
Œ3ˆu‹:ƒœ#˜e›*Ò$ÜÐBÓCÐCØ#€DØ€FÙ
Ø�I‰I�a‹LˆØˆ<Ø‰DÜ˜œcÔ"Ø�D˜1‘HŠ}Ü˜T 1 q¡5¨!Ó,‘à˜4‘x�Ü˜T 1 t¡8¨TÓ2‘ä˜d¤EÔ*Ü$Ð%HÓIÐIØ�D—I‘IŠ~Ü˜TŸZ™Z¨¯©°T·Y±YÑ)>ÀÇ	Á	ÓJ‘à—‘˜dÔ#Ø�ò# ô& �$œÔØ�‰”e˜D $¨¡(¨AÓ.Õ/Ü	�Dœ%Ô	 Ø�‰�dÔà€FØò 	:ˆÜˆq‹6�QŠ;à�M‰M˜QŸW™W˜IÕ'Ø�V‰V�qŠ[à�M‰M˜QŸW™W˜I Q q§v¡v hÐ/Õ0ð �M‰M˜QŸW™W˜I Q q§v¡v h¨a°·±¨xÐ8Õ9ð	:ð �8‰8�FÓÐr   c                ó@  — | j                  «       }t        |j                  «       «      D �cg c]  }t        j                  |«      ‘Œ }}t        j
                  j                  ||«       |D �cg c]b  }|d d j                  t        j                  «      j                  «       |dd  j                  t        j                  «      j                  «       f‘Œd }}t        t        «      }t        |«      D ]  \  }\  }	}
||	|
f   j                  |«       Œ  |dfS c c}w c c}w )Né   z(Seed, Offset))Ú	get_staterG   ÚsizeÚtorchÚ
empty_likeÚdistributedr   ÚviewÚuint64Úitemr   rS   r8   Úadd)Ú	generatorr&   Úlocal_stateÚ_Ú
all_statesÚstateÚseeds_offsetsÚseed_offset_ranksr!   ÚseedÚoffsets              r    Ú_check_philox_rng_syncru      s  € ð ×%Ñ%Ó'€KÜ9>¸u¿z¹z»|Ó9LÖM°A”%×"Ñ" ;Õ/ÐM€JÐMÜ	×Ñ× Ñ  ¨[Ô9ð  öàð 
ˆr�ˆ�‰œŸ™Ó	%×	*Ñ	*Ó	,¨e°A°B¨i¯n©n¼U¿\¹\Ó.J×.OÑ.OÓ.QÒRð€Mð ô $¤CÓ(ÐÜ )¨-Ó 8ò 4Ñˆ‰nˆt�VØ˜4 ˜.Ñ)×-Ñ-¨dÕ3ð4àÐ.Ð.Ð.ùò Nùòs   ¬DÁ.A'Dc                ó”  — | j                  «       }t        |j                  «       «      D �cg c]  }t        j                  |«      ‘Œ }}t        j
                  j                  ||«       t        t        «      }t        |«      D ]:  \  }}|t        j                  |«      j                  «          j                  |«       Œ< |dfS c c}w )NzGenerator state hash)rc   rG   rd   re   rf   rg   r   r   rS   r8   Úhash_tensorrj   rk   )rl   r&   Ústate_tensorrn   Úall_state_tensorsÚstate_ranksr!   s          r    Ú_check_cpu_rng_syncr{   0  s¶   € ð ×&Ñ&Ó(€LÜAFÀuÇzÁzÃ|ÓATÖU¸Aœ×)Ñ)¨,Õ7ÐUÐÐUÜ	×Ñ× Ñ Ð!2°LÔAÜœcÓ"€KÜ'Ð(9Ó:ò FÑˆˆlð 	”E×%Ñ% lÓ3×8Ñ8Ó:Ñ;×?Ñ?ÀÕEð	Fð
 Ð.Ð.Ð.ùò Vs   ¬Cc                óÚ   — | j                   j                  dk(  rt        | |«      S | j                   j                  dk(  rt        | |«      S t	        d| j                   j                  › �«      ‚)NÚcudaÚcpuzUnsupported generator device: )ÚdevicerA   ru   r{   ÚNotImplementedError)rl   r&   s     r    Ú_check_rng_sync_internalr�   @  sj   € ð ×Ñ×Ñ Ò&Ü% i°Ó7Ð7Ø	×	Ñ	×	Ñ	 %Ò	'Ü" 9¨eÓ4Ð4ä!Ø,¨Y×-=Ñ-=×-BÑ-BÐ,CÐDó
ð 	
r   c                ó`  — d| › d�g}|j                  «       D ��cg c]  \  }}t        |«      t        |«      g‘Œ }}}t        j                  j                  d«      rddlm}  |||¬«      S dj                  |D �cg c]  }t        |«      ‘Œ c}«      }t        |› d|› �«      S c c}}w c c}w )NÚRanksz valuesÚtabulater   )r„   )Úheadersú
)Úitemsr`   ÚstrÚ	importlibÚutilÚ	find_specr„   rZ   )	ÚtagÚvalue_ranksr…   Úvaluer[   Úrank_valuesr„   ÚrowÚrow_strs	            r    Ú_desync_table_strr’   M  s©   € Ø˜3˜%˜w˜Ð(€GàBM×BSÑBSÓBU÷Ù2>°%¸Ô	˜%Ó	 ¤# e£*Ò-ð€Kñ ô ‡~�~×Ñ 
Ô+Ý%á˜¨WÔ5Ð5Ø�i‰i¨[Ö9 cœ˜S�Ò9Ó:€GÜ�'�˜"˜W˜IÐ&Ó'Ð'ùóùò :s   › B%Á<B+c                óŒ   — t        | |«      \  }}d }t        |«      dkD  r$dt        ||«      › �}t        j	                  |«       |S )Nr'   zGenerator desync detected:
)r�   r-   r’   ÚloggerÚerror)rl   r&   r�   Úvalue_headerÚlog_strs        r    Ú_check_rng_syncr˜   Z  sL   € Ü 8¸ÀEÓ JÑ€K�Ø€GÜ
ˆ;Ó˜!ÒØ0Ô1BÀ<ÐQ\Ó1]Ð0^Ð_ˆÜ�‰�WÔØ€Nr   )r/   úT | Callable[[], T]r   r   r   r   r!   rV   r"   údist.ProcessGroup | NoneÚreturnr   )NN)r/   r™   r   r   r"   rš   r›   zlist[T])
r"   údist.ProcessGrouprI   z	list[Any]rJ   r   rK   zCallable[[Any, Any], bool]r›   ÚNone)r[   zIterable[int]r›   rˆ   )rl   útorch.Generatorr&   rœ   r›   ztuple[dict[Any, set], str])rŒ   rˆ   r�   zdict[Any, set[int]]r›   rˆ   )rl   rž   r&   rœ   r›   r   )&Ú__doc__Ú
__future__r   r‰   ÚloggingÚcollectionsr   Údataclassesr   Útypingr   r   r   r	   r
   Úcollections.abcr   r   re   Útorch.distributedrg   r+   r   r   Ú	getLoggerr   r”   r   r   r   r   r   r`   ru   r{   r�   r’   r˜   r   r   r    ú<module>r¨      s¨  ðòõ #ã Û Ý #Ý !ß =Õ =ñ ß2ã Ý  ò€ˆó ð 
ˆ×	Ñ	˜8Ó	$€áˆCƒL€ð ô'�'˜!‘*ó 'ó ð'ð Ø!ØØ#'ñM%Ø#ðM%ð ðM%ð ð	M%ð
 ðM%ð 	!ðM%ð óM%ðd "Ø#'ðI"Ø#ðI"àðI"ð 	!ðI"ð ó	I"ñl 0Oð"Øð"ð ð"ð
 
ð"ð -ð"ð 
ó"óJ+ð\/Øð/Ø'8ð/àó/ð /Øð/Ø'8ð/àó/ð 

Øð

Ø'8ð

àó

ó
(ôr   