§
    ŠŠtjÐ4  ã                   óÂ   — d dl Z d dlmZmZ d dlmZ d dlmZmZ d dl	Z	d dl
mZ g d¢Z G d„ d¦  «        Z G d„ d	e¦  «        Z G d
„ de¦  «        Z G d„ d¦  «        ZdS )é    N)ÚABCÚabstractmethod)ÚTracebackType)ÚAnyÚ
NamedTuple)ÚJoinHookÚJoinableÚJoinc                   ó*   — e Zd ZdZdd„Zdeddfd„ZdS )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                 ó   — dS )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    ú_/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/distributed/algorithms/join.pyÚ	main_hookzJoinHook.main_hook   ó   € € € ó    Úis_last_joinerc                 ó   — dS )aK  
        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_hookzJoinHook.post_hook    r   r   ©r   N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Úboolr   r   r   r   r   r      sT   € € € € € ð	ð 	ðð ð ð ð	¨ð 	°ð 	ð 	ð 	ð 	ð 	ð 	r   r   c                   ó²   ‡ — e Zd ZdZedˆ fd„¦   «         Zedefd„¦   «         Zeede	j
        fd„¦   «         ¦   «         Zeedefd„¦   «         ¦   «         Zˆ xZS )	r	   a_  
    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   Nc                 ó„   •— t          ¦   «                              ¦   «          t                               ¦   «         | _        d S ©N)ÚsuperÚ__init__Ú_JoinConfigÚconstruct_disabled_join_configÚ_join_config)r   Ú	__class__s    €r   r"   zJoinable.__init__7   s3   ø€ å‰Œ×ÒÑÔÐÝ'×FÒFÑHÔHˆÔÐÐr   c                 ó   — dS )aŽ  
        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_hookzJoinable.join_hook<   s	   € ð 	ˆr   c                 ó   — dS )zeReturn the device from which to perform collective communications needed by the join context manager.Nr   r   s    r   Újoin_devicezJoinable.join_deviceI   ó	   € ð 	ˆr   c                 ó   — dS )zfReturns the process group for the collective communications needed by the join context manager itself.Nr   r   s    r   Újoin_process_groupzJoinable.join_process_groupO   r,   r   r   )r   r   r   r   r   r"   r   r)   ÚpropertyÚtorchÚdevicer+   r   r.   Ú__classcell__)r&   s   @r   r	   r	   ,   sä   ø€ € € € € ðð ð ðIð Ið Ið Ið Iñ „^ðIð ð
 Xð 
ð 
ð 
ñ „^ð
ð Øð˜Uœ\ð ð ð ñ „^ñ „Xðð Øð Cð ð ð ñ „^ñ „Xðð ð ð ð r   r	   c                   óH   — e Zd ZU dZeed<   eed<   eed<   ed„ ¦   «         ZdS )r#   zdThis includes all fields needed from a :class:`Joinable` instance for the join context manager side.ÚenableÚthrow_on_early_terminationÚis_first_joinablec                  ó&   — t          ddd¬¦  «        S )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©r4   r5   r6   )r#   r   r   r   r$   z*_JoinConfig.construct_disabled_join_config]   s"   € õ Ø°UÈeð
ñ 
ô 
ð 	
r   N)r   r   r   r   r   Ú__annotations__Ústaticmethodr$   r   r   r   r#   r#   V   sV   € € € € € € ØoÐoà€L€L�LØ $Ð$Ð$Ñ$ØÐÐÑàð
ð 
ñ „\ð
ð 
ð 
r   r#   c                   ó¨   — e Zd ZdZ	 	 ddee         dedefd„Zdd
„Zdd„Z	d„ Z
dee         d	z  ded	z  ded	z  fd„Zd„ Zd„ Zedefd„¦   «         Zd	S )r
   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
    TFÚ	joinablesr4   r5   c                 óö   ‡— t          |¦  «        dk    rt          d¦  «        ‚|| _        ˆfd„| j        D ¦   «         | _        || _        || _        |                      ¦   «          |                      ¦   «          d S )Nr   z7The join context manager requires at least one joinablec                 ó*   •— g | ]} |j         d i ‰¤Ž‘ŒS )r   )r)   )Ú.0Újoinabler(   s     €r   ú
<listcomp>z!Join.__init__.<locals>.<listcomp>°   s:   ø€ ð 
ð 
ð 
Ø-5ÐˆHÔÐ(Ð( Ð(Ð(ð
ð 
ð 
r   )ÚlenÚ
ValueErrorÚ
_joinablesÚ_join_hooksÚ_enableÚ_throw_on_early_terminationÚ_set_joinable_configsÚ_extract_dist_info)r   r<   r4   r5   r(   s       `r   r"   zJoin.__init__¦   s”   ø€ õ ˆy‰>Œ>˜QÒÐÝÐVÑWÔWÐWØ#ˆŒð
ð 
ð 
ð 
Ø9=¼ð
ñ 
ô 
ˆÔð ˆŒØ+EˆÔ(Ø×"Ò"Ñ$Ô$Ð$Ø×ÒÑ!Ô!Ð!Ð!Ð!r   r   Nc                 ó¢   — t          | j        ¦  «        dk    rt          ‚d}| j        D ]%}t          | j        | j        |¬¦  «        |_        d}Œ&dS )zESet the :class:`_JoinConfig` of each participating :class:`Joinable`.r   Tr8   FN)rB   rD   ÚAssertionErrorr#   rF   rG   r%   )r   r6   r@   s      r   rH   zJoin._set_joinable_configs¸   sn   € åˆtŒÑÔ 1Ò$Ð$Ý Ð Ø ÐØœð 	&ð 	&ˆHÝ$/Ø”|Ø+/Ô+KØ"3ð%ñ %ô %ˆHÔ!ð
 !&ÐÐð	&ð 	&r   c                 óÔ   — d}d}| j         D ]/}|€|j        }n||j        k    rt          d¦  «        ‚|€|j        }Œ0|| _        t          j        | j        ¦  «        | _        || _        dS )aÃ  
        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)	rD   r.   rC   r+   Ú_process_groupÚdistÚget_rankÚ_rankÚ_device)r   Úprocess_groupr1   r@   s       r   rI   zJoin._extract_dist_infoÅ   s‰   € ð ˆØˆàœð 	.ð 	.ˆHØÐ$Ø (Ô ;��Ø (Ô"=Ò=Ð=Ý ØMñô ð ð ˆ~Ø!Ô-�øØ+ˆÔÝ”] 4Ô#6Ñ7Ô7ˆŒ
ØˆŒˆˆr   c                 ó   — d S r    r   r   s    r   Ú	__enter__zJoin.__enter__ä   r   r   ÚtypeÚvalueÚ	tracebackc           	      óª  — | j         r|rdS d}d}d}d}t          j        d¦  «         |sŠ||k    r%t          j        d|› d| j        › d	|› d
�d¬¦  «         |                      ¦   «         }|dk    rd}n@| j        r|                      ¦   «          | j        D ]}	|	 	                    ¦   «          Œd}|dz  }|¯Š| j        D ]}	|	 
                    |¦  «         ŒdS )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é   )rF   ÚwarningsÚsimplefilterÚwarnrP   Ú_get_num_nonjoined_procsrG   Ú_notify_procs_to_terminaterE   r   r   )
r   rU   rV   rW   Úall_procs_joinedr   ÚiÚWARN_THRESHOLDÚnum_nonjoined_procsr)   s
             r   Ú__exit__zJoin.__exit__æ   s_  € ð Œ|ð 	˜tð 	ØˆFà ÐØˆàˆØˆÝÔ˜fÑ%Ô%Ð%à"ð 	Ø�>Ò!Ð!Ý”ð3Ø%ð3ð 3à”zð3ð 3à1?ð3ð 3ð 3ð  !ðñ ô ð ð #'×"?Ò"?Ñ"AÔ"AÐØ" aÒ'Ð'Ø#'Ð Ð àÔ3ð 6Ø×3Ò3Ñ5Ô5Ð5ð "&Ô!1ð *ð *�IØ×'Ò'Ñ)Ô)Ð)Ð)à!&�Ø�Q‘�ð1 #ð 	ð6 Ô)ð 	0ð 	0ˆIØ×Ò Ñ/Ô/Ð/Ð/ð	0ð 	0r   c                 ó–   — t          j        d| j        ¬¦  «        }t          j        || j        ¬¦  «         |                     ¦   «         S )zaReturn the number of non-joined processes by shadowing an all-reduce in the non-joined processes.r\   ©r1   ©Úgroup)r0   ÚzerosrQ   rN   Ú
all_reducerM   Úitem)r   re   s     r   r`   zJoin._get_num_nonjoined_procs  sD   € å#œk¨!°D´LÐAÑAÔAÐÝŒÐ+°4Ô3FÐGÑGÔGÐGØ"×'Ò'Ñ)Ô)Ð)r   c                 óž   — t          j        d| j        ¬¦  «        }t          j        || j        ¬¦  «         t          d| j        › d�¦  «        ‚)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.
        r\   rh   ri   zRank z exhausted all inputs.)r0   ÚonesrQ   rN   rl   rM   ÚRuntimeErrorrP   )r   ro   s     r   ra   zJoin._notify_procs_to_terminate!  sN   € õ
 Œz˜! D¤LÐ1Ñ1Ô1ˆÝŒ˜ DÔ$7Ð8Ñ8Ô8Ð8ÝÐE 4¤:ÐEÐEÐEÑFÔFÐFr   r@   c                 óº  — t          | d¦  «        s t          dt          | ¦  «        › d�¦  «        ‚| j        }|j        r|j        sdS | j        }| j        }t          j	        d|¬¦  «        }t          j        ||d¬¦  «        }|j        rQt          j        d|¬¦  «        }t          j        ||¬	¦  «         |                     ¦   «         }|rt          d
¦  «        ‚|S )aH  
        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`` constructorNr\   rh   T)rj   Úasync_opri   zLDetected at least one rank that exhausted inputs. Throwing across all ranks.)ÚhasattrrK   rU   r%   r6   r4   r+   r.   r0   ro   rN   rl   r5   rk   rm   rp   )r@   Újoin_configr1   rR   ro   Úworkrk   Úshould_throws           r   Únotify_join_contextzJoin.notify_join_context*  s  € õ4 �x Ñ0Ô0ð 	Ý ð+¥$ x¡.¤.ð +ð +ð +ñô ð ð
 Ô+ˆàÔ,ð 	°KÔ4Fð 	Ø�4àÔ%ˆØ Ô3ˆõ Œz˜! FÐ+Ñ+Ô+ˆÝŒ˜t¨=À4ÐHÑHÔHˆàÔ1ð 		å”K ¨&Ð1Ñ1Ô1ˆEÝŒO˜E¨Ð7Ñ7Ô7Ð7Ø Ÿ:š:™<œ<ˆLØð Ý"ð1ñô ð ð ˆr   )TFr   )r   r   r   r   Úlistr	   r   r"   rH   rI   rT   rU   ÚBaseExceptionr   rf   r`   ra   r:   rw   r   r   r   r
   r
   h   s   € € € € € ð;ð ;ð@ Ø+0ð	"ð "à˜”>ð"ð ð"ð %)ð	"ð "ð "ð "ð$&ð &ð &ð &ðð ð ð ð> ÐÐð30à�=Ô! DÑ(ð30ð ˜tÑ#ð30ð ! 4Ñ'ð	30ð 30ð 30ð 30ðj*ð *ð *ðGð Gð Gð ð5 hð 5ð 5ð 5ñ „\ð5ð 5ð 5r   r
   )r]   Úabcr   r   Útypesr   Útypingr   r   r0   Útorch.distributedÚdistributedrN   Ú__all__r   r	   r#   r
   r   r   r   ú<module>r€      s1  ðà €€€Ø #Ð #Ð #Ð #Ð #Ð #Ð #Ð #Ø Ð Ð Ð Ð Ð Ø "Ð "Ð "Ð "Ð "Ð "Ð "Ð "à €€€Ø  Ð  Ð  Ð  Ð  Ð  ð +Ð
*Ð
*€ðð ð ð ð ñ ô ð ð<'ð 'ð 'ð 'ð 'ˆsñ 'ô 'ð 'ðT
ð 
ð 
ð 
ð 
�*ñ 
ô 
ð 
ð$xð xð xð xð xñ xô xð xð xð xr   