ó
    Eñi…4  ã                   ó¨   • S SK r S SKJrJr  S SKJr  S SKJrJr  S SK	r	S SK
Jr  / SQr " S S5      r " S S	\5      r " S
 S\5      r " S S5      rg)é    N)ÚABCÚabstractmethod)ÚTracebackType)ÚAnyÚ
NamedTuple)ÚJoinHookÚJoinableÚJoinc                   ó4   • \ rS rSrSrS	S jrS\SS4S jrSrg)
r   é   a¯  
This defines a join hook, which provides two entry points in the join context manager.

Entry points : a main hook, which is called repeatedly while there exists a non-joined
process, and a post-hook, which is called once all processes have joined.

To implement a join hook for the generic join context manager, define a
class that inherits from :class:`JoinHook` and override ``main_hook()`` and
``post_hook()`` as appropriate.
ÚreturnNc                 ó   • g)zÆCall this hook while there exists a non-joined process to shadow collective communications in a training iteration.

Training iteration i.e., in one forward pass, backward pass, and optimizer step.
N© ©Úselfs    Ú^/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/distributed/algorithms/join.pyÚ	main_hookÚJoinHook.main_hook   ó   � ó    Úis_last_joinerc                 ó   • g)a  
Call hook after all processes have joined.

It is passed an additional ``bool`` argument ``is_last_joiner``, which indicates if the rank is one of the last to join.

Arguments:
    is_last_joiner (bool): ``True`` if the rank is one of the last to
        join; ``False`` otherwise.
Nr   )r   r   s     r   Ú	post_hookÚJoinHook.post_hook    r   r   r   ©r   N)	Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   Úboolr   Ú__static_attributes__r   r   r   r   r      s   † ñ	ôð	¨ð 	°÷ 	r   r   c                   óª   ^ • \ rS rSrSr\S	U 4S jj5       r\S\4S j5       r\	\S\
R                  4S j5       5       r\	\S\4S j5       5       rSrU =r$ )
r	   é,   aC  
This defines an abstract base class for joinable classes.

A joinable class
(inheriting from :class:`Joinable`) should implement :meth:`join_hook`,
which returns a :class:`JoinHook` instance, in addition to
:meth:`join_device` and :meth:`join_process_group` that return device and
process group information, respectively.
r   c                 óT   >• [         TU ]  5         [        R                  5       U l        g ©N)ÚsuperÚ__init__Ú_JoinConfigÚconstruct_disabled_join_configÚ_join_config)r   Ú	__class__s    €r   r(   ÚJoinable.__init__7   s   ø€ ä‰ÑÔÜ'×FÑFÓHˆÕr   c                 ó   • g)aV  
Return a :class:`JoinHook` instance for the given :class:`Joinable`.

Arguments:
    kwargs (dict): a :class:`dict` containing any keyword arguments
        to modify the behavior of the join hook at run time; all
        :class:`Joinable` instances sharing the same join context
        manager are forwarded the same value for ``kwargs``.
Nr   )r   Úkwargss     r   Ú	join_hookÚJoinable.join_hook<   s   € ð 	r   c                 ó   • g)zeReturn the device from which to perform collective communications needed by the join context manager.Nr   r   s    r   Újoin_deviceÚJoinable.join_deviceI   ó   € ð 	r   c                 ó   • g)zfReturns the process group for the collective communications needed by the join context manager itself.Nr   r   s    r   Újoin_process_groupÚJoinable.join_process_groupO   r5   r   )r+   r   )r   r   r   r   r    r   r(   r   r0   ÚpropertyÚtorchÚdevicer3   r   r7   r"   Ú__classcell__)r,   s   @r   r	   r	   ,   sƒ   ø† ñð öIó ðIð ð
 Xó 
ó ð
ð Øð˜UŸ\™\ó ó ó ðð Øð Có ó ó ör   r	   c                   óH   • \ rS rSr% Sr\\S'   \\S'   \\S'   \S 5       rSr	g)	r)   éV   zdThis includes all fields needed from a :class:`Joinable` instance for the join context manager side.ÚenableÚthrow_on_early_terminationÚis_first_joinablec                  ó   • [        SSSS9$ )z”Return a :class:`_JoinConfig` instance indicating that join-related logic should be disabled.

e.g. if the caller is not in a join context manager.
F©r?   r@   rA   )r)   r   r   r   r*   Ú*_JoinConfig.construct_disabled_join_config]   s   € ô Ø°UÈeñ
ð 	
r   r   N)
r   r   r   r   r    r!   Ú__annotations__Ústaticmethodr*   r"   r   r   r   r)   r)   V   s(   ‡ ÙoàƒLØ $Ó$ØÓàñ
ó ó
r   r)   c                   ó¨   • \ rS rSrSr  SS\\   S\S\4S jjrSS jr	SS	 jr
S
 rS\\   S-  S\S-  S\S-  4S jrS rS r\S\4S j5       rSrg)r
   éh   a
  
This class defines the generic join context manager, which allows custom hooks to be called after a process joins.

These hooks should shadow the
collective communications of non-joined processes to prevent hanging and
erroring and to ensure algorithmic correctness. Refer to :class:`JoinHook`
for details about the hook definition.

.. warning::
    The context manager requires each participating :class:`Joinable` to
    call the method :meth:`notify_join_context()` before its own per-
    iteration collective communications to ensure correctness.

.. warning::
    The context manager requires that all ``process_group`` attributes in
    the :class:`JoinHook` objects are the same. If there are multiple
    :class:`JoinHook` objects, then the ``device`` of the first is used.
    The process group and device information is used for checking for non-
    joined processes and for notifying processes to throw an exception if
    ``throw_on_early_termination`` is enabled, both of which using an all-
    reduce.

Arguments:
    joinables (List[Joinable]): a list of the participating
        :class:`Joinable` s; their hooks are iterated over in the given
        order.

    enable (bool): a flag enabling uneven input detection; setting to
        ``False`` disables the context manager's functionality and should
        only be set when the user knows the inputs will not be uneven
        (default: ``True``).

    throw_on_early_termination (bool): a flag controlling whether to throw an
        exception upon detecting uneven inputs (default: ``False``).

Example::

    >>> import os
    >>> import torch
    >>> import torch.distributed as dist
    >>> import torch.multiprocessing as mp
    >>> # xdoctest: +SKIP
    >>> import torch.nn.parallel.DistributedDataParallel as DDP
    >>> import torch.distributed.optim.ZeroRedundancyOptimizer as ZeRO
    >>> from torch.distributed.algorithms.join import Join
    >>>
    >>> # On each spawned worker
    >>> def worker(rank):
    >>>     dist.init_process_group("nccl", rank=rank, world_size=2)
    >>>     model = DDP(torch.nn.Linear(1, 1).to(rank), device_ids=[rank])
    >>>     optim = ZeRO(model.parameters(), torch.optim.Adam, lr=0.01)
    >>>     # Rank 1 gets one more input than rank 0
    >>>     inputs = [torch.tensor([1.]).to(rank) for _ in range(10 + rank)]
    >>>     with Join([model, optim]):
    >>>         for input in inputs:
    >>>             loss = model(input).sum()
    >>>             loss.backward()
    >>>             optim.step()
    >>>     # All ranks reach here without hanging/erroring
Ú	joinablesr?   r@   c                 ó  • [        U5      S:X  a  [        S5      eXl        U R                   Vs/ s H  oUR                  " S0 UD6PM     snU l        X l        X0l        U R                  5         U R                  5         g s  snf )Nr   z7The join context manager requires at least one joinabler   )	ÚlenÚ
ValueErrorÚ
_joinablesr0   Ú_join_hooksÚ_enableÚ_throw_on_early_terminationÚ_set_joinable_configsÚ_extract_dist_info)r   rI   r?   r@   r/   Újoinables         r   r(   ÚJoin.__init__¦   sw   € ô ˆy‹>˜QÓÜÐVÓWÐWØ#Œà9=¿ºó
Ú9H¨X×ÒÑ( Ô(¹ñ
ˆÔð ŒØ+EÔ(Ø×"Ñ"Ô$Ø×ÑÕ!ùò
s   ¯A?Nc                 ó°   • [        U R                  5      S:”  d   eSnU R                   H)  n[        U R                  U R                  US9Ul        SnM+     g)zESet the :class:`_JoinConfig` of each participating :class:`Joinable`.r   TrC   FN)rK   rM   r)   rO   rP   r+   )r   rA   rS   s      r   rQ   ÚJoin._set_joinable_configs¸   sU   € ä�4—?‘?Ó# aÓ'Ð'Ð'Ø ÐØŸœˆHÜ$/Ø—|‘|Ø+/×+KÑ+KØ"3ñ%ˆHÔ!ð
 !&Òò (r   c                 ó
  • SnSnU R                    H>  nUc  UR                  nOXR                  :w  a  [        S5      eUb  M2  UR                  nM@     Xl        [
        R                  " U R                  5      U l        X l        g)as  
Extract the process group and device information from the joinables.

If there are multiple joinables, then the context manager uses the
first specified device.

Preconditions:
    ``self._joinables`` is not ``None`` and is non-empty.

Raises:
    ValueError
        If there are multiple conflicting ``process_group`` attributes
        among the ``Joinable`` objects.
Nz7Using join context manager with multiple process groups)	rM   r7   rL   r3   Ú_process_groupÚdistÚget_rankÚ_rankÚ_device)r   Úprocess_groupr;   rS   s       r   rR   ÚJoin._extract_dist_infoÄ   s~   € ð ˆØˆàŸœˆHØÑ$Ø (× ;Ñ ;‘Ø×"=Ñ"=Ó=Ü ØMóð ð ‹~Ø!×-Ñ-’ñ (ð ,ÔÜ—]’] 4×#6Ñ#6Ó7ˆŒ
Ø�r   c                 ó   • g r&   r   r   s    r   Ú	__enter__ÚJoin.__enter__ã   s   € ˜r   ÚtypeÚvalueÚ	tracebackc           	      óþ  • U R                   (       a  U(       a  gSnSnSnSn[        R                  " S5        U(       d›  Xg:”  a)  [        R                  " SU SU R                   S	U S
3SS9  U R                  5       nUS:X  a  SnOKU R                  (       a  U R                  5         U R                   H  n	U	R                  5         M     SnUS-  nU(       d  M›  U R                   H  n	U	R                  U5        M     g)zŸ
Repeatedly runs the main hooks until all processes join; then, runs the post-hooks.

Raises:
    RuntimeError
        If ``throw_on_early_termination=True``.
NFTr   iè  Úoncez+Detected uneven input skew of greater than z. This means that rank z has at least zz fewer inputs than other currently-active ranks. This level of skew could lead to performance degradation during training.é   )Ú
stacklevelé   )rO   ÚwarningsÚsimplefilterÚwarnr[   Ú_get_num_nonjoined_procsrP   Ú_notify_procs_to_terminaterN   r   r   )
r   rb   rc   rd   Úall_procs_joinedr   ÚiÚWARN_THRESHOLDÚnum_nonjoined_procsr0   s
             r   Ú__exit__ÚJoin.__exit__å   sþ   € ð �|�|žtØà ÐØˆàˆØˆÜ×Ò˜fÔ%æ"ØÓ!Ü—’ØAØ%Ð&Ð&=Ø—z‘z�l .°Ð0@ð A3ð3ð  !òð #'×"?Ñ"?Ó"AÐØ" aÓ'Ø#'Ñ à×3×3Ø×3Ñ3Ô5ð "&×!1Ô!1�IØ×'Ñ'Ö)ñ "2ð "'�Ø�Q‘�÷1 #Ð"ð6 ×)Ô)ˆIØ×Ñ Ö/ò *r   c                 óœ   • [         R                  " SU R                  S9n[        R                  " XR
                  S9  UR                  5       $ )zaReturn the number of non-joined processes by shadowing an all-reduce in the non-joined processes.ri   ©r;   ©Úgroup)r:   Úzerosr\   rY   Ú
all_reducerX   Úitem)r   rr   s     r   rm   ÚJoin._get_num_nonjoined_procs  s9   € ä#Ÿkšk¨!°D·L±LÑAÐÜ�ŠÐ+×3FÑ3FÒGØ"×'Ñ'Ó)Ð)r   c                 ó®   • [         R                  " SU R                  S9n[        R                  " XR
                  S9  [        SU R                   S35      e)z¡Schedule an all-reduce to notify non-joined processes to terminate.

Also raise a ``RuntimeError`` indicating that the current process has exhausted its inputs.
ri   rv   rw   zRank z exhausted all inputs.)r:   Úonesr\   rY   rz   rX   ÚRuntimeErrorr[   )r   r~   s     r   rn   ÚJoin._notify_procs_to_terminate   sC   € ô
 �zŠz˜! D§L¡LÑ1ˆÜ�Š˜×$7Ñ$7Ò8Ü˜U 4§:¡: ,Ð.DÐEÓFÐFr   rS   c                 óê  • [        U S5      (       d   S[        U 5       S35       eU R                  nUR                  (       a  UR                  (       d  gU R
                  nU R                  n[        R                  " SUS9n[        R                  " XCSS9nUR                  (       aK  [        R                  " SUS9n[        R                  " XcS	9  UR                  5       nU(       a  [        S
5      eU$ )a¸  
Notifies the join context manager that the calling process has not yet joined.

Then, if ``throw_on_early_termination=True``, checks if uneven inputs have been detected
(i.e. if one process has already joined) and throws an exception if so.

This method should be called from a :class:`Joinable` object before
its per-iteration collective communications. For example, this should
be called at the beginning of the forward pass in
:class:`DistributedDataParallel`.

Only the first :class:`Joinable` object passed into the context
manager performs the collective communications in this method, and
for the others, this method is vacuous.

Arguments:
    joinable (Joinable): the :class:`Joinable` object calling this
        method.

Returns:
    An async work handle for the all-reduce meant to notify the context
    manager that the process has not yet joined if ``joinable`` is the
    first one passed into the context manager; ``None`` otherwise.
r+   zCheck that the z/ constructor calls the ``Joinable`` constructorNri   rv   T)rx   Úasync_oprw   zLDetected at least one rank that exhausted inputs. Throwing across all ranks.)Úhasattrrb   r+   rA   r?   r3   r7   r:   r~   rY   rz   r@   ry   r{   r   )rS   Újoin_configr;   r]   r~   Úworkry   Úshould_throws           r   Únotify_join_contextÚJoin.notify_join_context)  sÙ   € ô4 �x ×0Ñ0ð 	
Øœd 8›nÐ-ð .'ð 'ó	
Ð0ð
 ×+Ñ+ˆà×,×,°K×4F×4FØà×%Ñ%ˆØ ×3Ñ3ˆô �zŠz˜! FÑ+ˆÜ�Š˜tÀ4ÑHˆà×1×1ä—K’K ¨&Ñ1ˆEÜ�OŠO˜EÒ7Ø Ÿ:™:›<ˆLÞÜ"ð1óð ð ˆr   )r\   rO   rN   rM   rX   r[   rP   )TFr   )r   r   r   r   r    Úlistr	   r!   r(   rQ   rR   r`   rb   ÚBaseExceptionr   rs   rm   rn   rF   r‡   r"   r   r   r   r
   r
   h   s    † ñ;ð@ Ø+0ñ	"à˜‘>ð"ð ð"ð %)õ	"ô$
&ôò> ð30à�=Ñ! DÑ(ð30ð ˜tÑ#ð30ð ! 4Ñ'ô	30òj*òGð ð4 hó 4ó ó4r   r
   )rj   Úabcr   r   Útypesr   Útypingr   r   r:   Útorch.distributedÚdistributedrY   Ú__all__r   r	   r)   r
   r   r   r   Ú<module>r‘      sM   ðã ß #Ý ß "ã Ý  ò +€÷ñ ô<'ˆsô 'ôT
�*ô 
÷$vò vr   