§
    ŠŠtjÂG  ã                   ó\  — 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 ddœd„Z	d„ Z
ej        fd„Zd ej        fd„Zd ej        fd	„Zej        ej        fd
„Zej        ej        fd„Zej        fd„Zej        fd„Zej        fd„Zddej        fd„Zej        ej        fd„Z G d„ de¦  «        Z G d„ de¦  «        Z G d„ de¦  «        Z G d„ de¦  «        Z G d„ de¦  «        Z G d„ de¦  «        Z G d„ de¦  «        Z G d„ d e¦  «        Z G d!„ d"e¦  «        Z G d#„ d$e¦  «        Z dS )%é    N)ÚFunction)ÚgroupÚReduceOp©Ú
suggestionc                óB   — d| › d�}|r	|d|› d�z  }t          |¦  «        ‚)Nú torch.distributed.nn.functional.z& is not supported under torch.compile.z Use ú	 instead.)ÚRuntimeError)Únamer   Úmsgs      ú]/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/distributed/nn/functional.pyÚ_not_supported_under_compiler      sC   € àW¨4ÐWÐWÐWð ð ð -ØÐ,�zÐ,Ð,Ð,Ñ,ˆÝ
�sÑ
Ô
Ðó    c                 óL   — t          j        d| › d|› d�t          d¬¦  «         d S )Nr	   z is deprecated, use r
   é   )ÚcategoryÚ
stacklevel)ÚwarningsÚwarnÚFutureWarning)r   r   s     r   Ú_deprecatedr      sN   € Ý„Mð	%¨4ð 	%ð 	%Øð	%ð 	%ð 	%åØð	ñ ô ð ð ð r   c                 ó¸   — t           j                             ¦   «         rt          dd¬¦  «         t	          dd¦  «         t
                               ||| ¦  «        S )a»  
    Broadcasts the tensor to the whole group.

    ``tensor`` must have the same number of elements in all processes
    participating in the collective.

    Arguments:
        tensor (Tensor): Data to be sent if ``src`` is the rank of current
            process.
        src (int): Source rank.
        group (ProcessGroup, optional): The process group to work on.

    Returns:
        Tensor: Received tensor from the broadcast op.

    Ú	broadcastz3torch.distributed._functional_collectives.broadcastr   )ÚtorchÚcompilerÚis_compilingr   r   Ú
_BroadcastÚapply)ÚtensorÚsrcr   s      r   r   r       sb   € õ" „~×"Ò"Ñ$Ô$ð 
Ý$ØØLð	
ñ 	
ô 	
ð 	
õ �ÐRÑSÔSÐSÝ×Ò˜C ¨Ñ/Ô/Ð/r   c                 ó”   — t           j                             ¦   «         rt          d¦  «         t                               ||| ¦  «        S )aT  
    Gathers a list of tensors in a single process.

    Arguments:
        tensor (Tensor): Input tensor.
        dst (int, optional): Destination rank (default is 0).
        group (ProcessGroup, optional): The process group to work on.

    Returns:
        tuple[Tensor]: List of appropriately-sized tensors with the gathered data.
    Úgather)r   r   r   r   Ú_Gatherr   )r    Údstr   s      r   r#   r#   :   s>   € õ „~×"Ò"Ñ$Ô$ð /Ý$ XÑ.Ô.Ð.Ý�=Š=˜˜e VÑ,Ô,Ð,r   c                 ó‚   — t           j                             ¦   «         rt          d¦  «         t	          j        ||g| ¢R Ž S )aö  
    Scatters a list of tensors to all processes in a group.

    Each process will receive exactly one tensor and store its data in the
    ``tensor`` argument.

    Arguments:
        tensors (list[Tensor]): List of tensors to scatter on the source rank.
            Receivers must pass ``None`.
        src (int, optional): Source rank (default is 0).
        group (ProcessGroup, optional): The process group to work on.

    Returns:
        Tensor: Output tensor from the scatter operation.

    Úscatter)r   r   r   r   Ú_Scatterr   )Útensorsr!   r   s      r   r'   r'   K   sB   € õ" „~×"Ò"Ñ$Ô$ð 0Ý$ YÑ/Ô/Ð/ÝŒ>˜#˜uÐ/ wÐ/Ð/Ð/Ð/r   c                 ó–   — t           j                             ¦   «         rt          d¦  «         t                               |||| ¦  «        S )a  
    Reduces the tensor data across all machines.

    Only the process with rank ``dst`` is going to receive the final result.

    Arguments:
        tensor (Tensor): Input of the collective.
        dst (int): Destination rank.
        op (optional): One of the values from
            ``torch.distributed.ReduceOp``
            enum.  Specifies an operation used for element-wise reductions.
        group (ProcessGroup, optional): The process group to work on.

    Returns:
        Tensor: Output of the collective.

    Úreduce)r   r   r   r   Ú_Reducer   )r    r%   Úopr   s       r   r+   r+   a   s@   € õ$ „~×"Ò"Ñ$Ô$ð /Ý$ XÑ.Ô.Ð.Ý�=Š=˜˜b %¨Ñ0Ô0Ð0r   c                 ó¨   — t           j                             ¦   «         rt          dd¬¦  «         t	          dd¦  «         t          j        ||| g|¢R Ž S )aõ  
    Reduces, then scatters a list of tensors to all processes in a group.

    Arguments:
        output (Tensor): Output tensor.
        input_list (list[Tensor]): List of tensors to reduce and scatter.
        op (optional): One of the values from
            ``torch.distributed.ReduceOp``
            enum.  Specifies an operation used for element-wise reductions.
        group (ProcessGroup, optional): The process group to work on.

    Returns:
        Tensor: Output of the collective.

    Úreduce_scatterz?torch.distributed._functional_collectives.reduce_scatter_singler   )r   r   r   r   r   Ú_Reduce_Scatterr   )ÚoutputÚ
input_listr-   r   s       r   r/   r/   x   sp   € õ  „~×"Ò"Ñ$Ô$ð 
Ý$ØØXð	
ñ 	
ô 	
ð 	
õ ØØIñô ð õ Ô   U¨FÐ@°ZÐ@Ð@Ð@Ð@r   c                 ó¶   — t           j                             ¦   «         rt          dd¬¦  «         t	          dd¦  «         t
                               || ¦  «        S )a  
    Gathers tensors from the whole group in a list.

    Arguments:
        tensor (Tensor): Tensor to be broadcast from current process.
        group (ProcessGroup, optional): The process group to work on.

    Returns:
        tuple([Tensor]): Output of the collective.

    Ú
all_gatherz;torch.distributed._functional_collectives.all_gather_singler   )r   r   r   r   r   Ú
_AllGatherr   )r    r   s     r   r4   r4   ”   si   € õ „~×"Ò"Ñ$Ô$ð 
Ý$ØØTð	
ñ 	
ô 	
ð 	
õ ØÐSñô ð õ ×Ò˜E 6Ñ*Ô*Ð*r   c                 ó”   — t           j                             ¦   «         rt          d¦  «         t                               | ||¦  «        S )aâ  
    Single tensor all gather. Gathers a single tensor from all ranks, and puts them in a single output tensor.

    Args:
        output_tensor (Tensor): Output tensor. It should contain
            correctly-sized tensors to be used for output of the collective.
        input_tensor (Tensor): Tensor to be broadcast from current process.
        group (ProcessGroup, optional): The process group to work on. If None,
            the default process group will be used.

    Examples:
        >>> # All tensors below are of torch.int64 dtype.
        >>> # We have 2 process groups, 2 ranks.
        >>> # xdoctest: +SKIP("incorrect want text")
        >>> output_tensor = torch.zeros(2, dtype=torch.int64)
        >>> output_tensor
        [tensor([0, 0])] # Rank 0 and 1
        >>> tensor = torch.arange(1, dtype=torch.int64) + 1 + rank
        >>> tensor
        tensor([1]) # Rank 0
        tensor([2]) # Rank 1
        >>> dist.all_gather_base(output_tensor, tensor)
        >>> output_tensor
        tensor([1,2]) # Rank 0
        tensor([1,2]) # Rank 1

    .. warning::
        `_all_gather_base` is experimental and subject to change.
        It is the caller's responsibility to ensure the output_tensor
        is correctly sized.

    Ú_all_gather_base)r   r   r   r   Ú_AllGatherBaser   )Úoutput_tensorÚinput_tensorr   s      r   r7   r7   «   sB   € õB „~×"Ò"Ñ$Ô$ð 9Ý$Ð%7Ñ8Ô8Ð8Ý×Ò ¨|¸UÑCÔCÐCr   c                 ó‚   — t           j                             ¦   «         rt          d¦  «         t	          j        || g|¢R Ž S )aÃ  
    Each process scatters list of input tensors to all processes in a group and return gathered list of tensors in output list.

    Arguments:
        output_tensor_list (list[Tensor]): list of tensors to gather one per rank.
        input_tensor_list (list[Tensor]): List of tensors to scatter one per rank.
        group (ProcessGroup, optional): The process group to work on.

    Returns:
        tuple([Tensor]): Output of the collective.

    Ú
all_to_all)r   r   r   r   Ú	_AlltoAllr   )Úoutput_tensor_listÚinput_tensor_listr   s      r   r<   r<   Ñ   sD   € õ „~×"Ò"Ñ$Ô$ð 3Ý$ \Ñ2Ô2Ð2ÝŒ?˜5Ð"4ÐIÐ7HÐIÐIÐIÐIr   c                 ó¼   — t           j                             ¦   «         rt          dd¬¦  «         t	          dd¦  «         t
                               || |||¦  «        S )a  
    Each process splits input tensor and then scatters the split list to all processes in a group.

    Then concatenate the received tensors from all the processes in the group and return single output tensor.

    Arguments:
        output (Tensor): Gathered concatenated output tensor.
        input (Tensor): Input tensor to scatter.
        output_split_sizes: (list[Int], optional): Output split sizes for dim 0
            if specified None or empty, dim 0 of ``output`` tensor must divide
            equally by ``world_size``.
        input_split_sizes: (list[Int], optional): Input split sizes for dim 0
            if specified None or empty, dim 0 of ``input`` tensor must divide
            equally by ``world_size``.

    Returns:
        Tensor: Output of the collective.

    Úall_to_all_singlez;torch.distributed._functional_collectives.all_to_all_singler   )r   r   r   r   r   Ú_AlltoAllSingler   )r1   ÚinputÚoutput_split_sizesÚinput_split_sizesr   s        r   rA   rA   ã   sx   € õ4 „~×"Ò"Ñ$Ô$ð 
Ý$ØØTð	
ñ 	
ô 	
ð 	
õ ØØEñô ð õ × Ò ØˆvÐ)Ð+<¸eñô ð r   c                 ó¸   — t           j                             ¦   «         rt          dd¬¦  «         t	          dd¦  «         t
                               ||| ¦  «        S )a&  
    Reduces the tensor data across all machines in such a way that all get the final result.

    After the call the returned tensor is going to be bitwise
    identical in all processes.

    Arguments:
        tensor (Tensor): Input of the collective.
        op (optional): One of the values from
            ``torch.distributed.ReduceOp``
            enum.  Specifies an operation used for element-wise reductions.
        group (ProcessGroup, optional): The process group to work on.

    Returns:
        Tensor: Output of the collective

    Ú
all_reducez4torch.distributed._functional_collectives.all_reducer   )r   r   r   r   r   Ú
_AllReducer   )r    r-   r   s      r   rG   rG     sb   € õ$ „~×"Ò"Ñ$Ô$ð 
Ý$ØØMð	
ñ 	
ô 	
ð 	
õ �ÐTÑUÔUÐUÝ×Ò˜B  vÑ.Ô.Ð.r   c                   ó:   — e Zd Zed„ ¦   «         Zed„ ¦   «         ZdS )r   c                 ó¬   — || _         || _        t          j        |¬¦  «        | _        |                     ¦   «         }t          j        |||¬¦  «         |S ©N©r   )r!   r   ÚdistÚget_rankÚrankÚcloner   )Úctxr!   r   r    s       r   Úforwardz_Broadcast.forward'  sQ   € ð ˆŒØˆŒ	Ý”= uÐ-Ñ-Ô-ˆŒð —’‘”ˆÝŒ�v˜s¨%Ð0Ñ0Ô0Ð0Øˆr   c                 ó¶   — t                                | j        t          j        | j        |¦  «        }| j        | j        k    r|                     ¦   «          d d |fS ©N)r,   r   r!   r   ÚSUMr   rO   Úzero_)rQ   Úgrad_outputÚgxs      r   Úbackwardz_Broadcast.backward3  sJ   € õ �]Š]˜3œ7¥H¤L°#´)¸[ÑIÔIˆØŒ7�c”hÒÐØ�HŠH‰JŒJˆJØ�d˜BÐÐr   N©Ú__name__Ú
__module__Ú__qualname__ÚstaticmethodrR   rY   © r   r   r   r   &  sH   € € € € € Øðð ñ „\ðð ð ð  ñ „\ð ð  ð  r   r   c                   ó:   — e Zd Zed„ ¦   «         Zed„ ¦   «         ZdS )r$   c                 óV  ‡— || _         || _        ˆfd„t          t          j        |¬¦  «        ¦  «        D ¦   «         }‰                     ¦   «         Št          j        |¬¦  «        |k    rt          j        ‰|||¬¦  «         nt          j        ‰d ||¬¦  «         t          |¦  «        S )Nc                 ó8   •— g | ]}t          j        ‰¦  «        ‘ŒS r_   )r   Ú
zeros_like)Ú.0Úir    s     €r   ú
<listcomp>z#_Gather.forward.<locals>.<listcomp>F  s1   ø€ ð 
ð 
ð 
Ø)*�EÔ˜VÑ$Ô$ð
ð 
ð 
r   rL   )	r%   r   ÚrangerM   Úget_world_sizeÚ
contiguousrN   r#   Útuple)rQ   r%   r   r    Útensor_lists      ` r   rR   z_Gather.forward=  s¿   ø€ ð ˆŒØˆŒ	ð

ð 
ð 
ð 
Ý.3µDÔ4GÈeÐ4TÑ4TÔ4TÑ.UÔ.Uð
ñ 
ô 
ˆð ×"Ò"Ñ$Ô$ˆÝŒ=˜uÐ%Ñ%Ô%¨Ò,Ð,ÝŒK˜ ¨S¸Ð>Ñ>Ô>Ð>Ð>åŒK˜  c°Ð7Ñ7Ô7Ð7Ý�[Ñ!Ô!Ð!r   c                 óD   — dt          j        | j        | j        g|¢R Ž fz   S ©N©NN)r(   r   r%   r   )rQ   Úgrad_outputss     r   rY   z_Gather.backwardQ  s(   € à�xœ~¨c¬g°s´yÐPÀ<ÐPÐPÐPÐRÑRÐRr   NrZ   r_   r   r   r$   r$   <  sM   € € € € € Øð"ð "ñ „\ð"ð$ ðSð Sñ „\ðSð Sð Sr   r$   c                   ó:   — e Zd Zed„ ¦   «         Zed„ ¦   «         ZdS )r(   c                 óJ  ‡— || _         || _        t          ˆfd„‰D ¦   «         ¦  «        st          ‚t	          j        ‰d         ¦  «        }t          j        |¬¦  «        |k    r&t          j        |t          ‰¦  «        ||¬¦  «         nt          j        |d ||¬¦  «         |S )Nc              3   óx   •K  — | ]4}|                      ¦   «         ‰d                                ¦   «         k    V — Œ5dS )r   N©Úsize)rd   Útr)   s     €r   ú	<genexpr>z#_Scatter.forward.<locals>.<genexpr>\  s>   øè è € ÐBÐB°Q�1—6’6‘8”8˜w qœzŸšÑ0Ô0Ò0ÐBÐBÐBÐBÐBÐBr   r   rL   )
r!   r   ÚallÚAssertionErrorr   rc   rM   rN   r'   Úlist)rQ   r!   r   r)   r1   s      ` r   rR   z_Scatter.forwardW  sª   ø€ ð ˆŒØˆŒ	ÝÐBÐBÐBÐB¸'ÐBÑBÔBÑBÔBð 	!Ý Ð ÝÔ! '¨!¤*Ñ-Ô-ˆÝŒ=˜uÐ%Ñ%Ô%¨Ò,Ð,ÝŒL˜¥ g¡¤°¸5ÐAÑAÔAÐAÐAåŒL˜  s°%Ð8Ñ8Ô8Ð8Øˆr   c                 óT   — dt                                | j        | j        |¦  «        z   S rm   )r$   r   r!   r   ©rQ   rW   s     r   rY   z_Scatter.backwarde  s#   € ð �gŸmšm¨C¬G°S´YÀÑLÔLÑLÐLr   NrZ   r_   r   r   r(   r(   V  sM   € € € € € Øð
ð 
ñ „\ð
ð ðMð Mñ „\ðMð Mð Mr   r(   c                   ó:   — e Zd Zed„ ¦   «         Zed„ ¦   «         ZdS )r,   c                 óz   — || _         || _        |                     ¦   «         }t          j        ||||¬¦  «         |S )N©r-   r   )r!   r   rP   rM   r+   )rQ   r!   r-   r   r    s        r   rR   z_Reduce.forwardl  s=   € ð ˆŒØˆŒ	Ø—’‘”ˆÝŒ�F˜C B¨eÐ4Ñ4Ô4Ð4Øˆr   c                 óV   — dt                                | j        | j        |¦  «        fz   S ©N)NNN)r   r   r!   r   r{   s     r   rY   z_Reduce.backwardu  s(   € ð "¥Z×%5Ò%5°c´g¸s¼yÈ+Ñ%VÔ%VÐ$XÑXÐXr   NrZ   r_   r   r   r,   r,   k  sM   € € € € € Øðð ñ „\ðð ðYð Yñ „\ðYð Yð Yr   r,   c                   ó:   — e Zd Zed„ ¦   «         Zed„ ¦   «         ZdS )r0   c                 ó¸   — || _         |                     ¦   «         }t          d„ |D ¦   «         ¦  «        }t          j        |t          |¦  «        ||¬¦  «         |S )Nc              3   ó>   K  — | ]}|                      ¦   «         V — Œd S rT   ©ri   ©rd   ru   s     r   rv   z*_Reduce_Scatter.forward.<locals>.<genexpr>‚  s*   è è € Ð!LÐ!L°Q !§,¢,¡.¤.Ð!LÐ!LÐ!LÐ!LÐ!LÐ!Lr   r~   )r   ri   rj   rM   r/   ry   )rQ   r-   r   r    r?   s        r   rR   z_Reduce_Scatter.forward|  sb   € ð ˆŒ	à×"Ò"Ñ$Ô$ˆÝ!Ð!LÐ!LÐ:KÐ!LÑ!LÔ!LÑLÔLÐÝÔ˜F¥DÐ):Ñ$;Ô$;ÀÈ%ÐPÑPÔPÐPØˆr   c                 óH   — dt                                | j        |¦  «        z   S r€   )r5   r   r   r{   s     r   rY   z_Reduce_Scatter.backward†  s!   € ð "¥J×$4Ò$4°S´YÀÑ$LÔ$LÑLÐLr   NrZ   r_   r   r   r0   r0   {  sM   € € € € € Øðð ñ „\ðð ðMð Mñ „\ðMð Mð Mr   r0   c                   ó:   — e Zd Zed„ ¦   «         Zed„ ¦   «         ZdS )r5   c                 óâ   ‡— ‰                      ¦   «         Š|| _        ˆfd„t          t          j        |¬¦  «        ¦  «        D ¦   «         }t          j        |‰|¬¦  «         t          |¦  «        S )Nc                 ó8   •— g | ]}t          j        ‰¦  «        ‘ŒS r_   ©r   Ú
empty_like)rd   Ú_r    s     €r   rf   z&_AllGather.forward.<locals>.<listcomp>”  s1   ø€ ð 
ð 
ð 
Ø)*�EÔ˜VÑ$Ô$ð
ð 
ð 
r   rL   )ri   r   rg   rM   rh   r4   rj   )rQ   r   r    Úout_tensor_lists     ` r   rR   z_AllGather.forward�  s‚   ø€ ð ×"Ò"Ñ$Ô$ˆàˆŒ	ð
ð 
ð 
ð 
Ý.3µDÔ4GÈeÐ4TÑ4TÔ4TÑ.UÔ.Uð
ñ 
ô 
ˆõ 	Œ˜¨°uÐ=Ñ=Ô=Ð=Ý�_Ñ%Ô%Ð%r   c                 óÊ  — t          j        | j        ¬¦  «        t           j        j        t           j        j        fv rXt          j        | j        ¬¦  «        }t          j        ||         ¦  «        }t          j
        t          j        | j        |g|¢R Ž }nLd„ |D ¦   «         }t          j
        | j        |g|¢R Ž }t          j        t          j        |¦  «        d¬¦  «        }d |fS )NrL   c                 ó6   — g | ]}t          j        |¦  «        ‘ŒS r_   rŠ   )rd   r    s     r   rf   z'_AllGather.backward.<locals>.<listcomp>¤  s#   € ÐOÐOÐO¸�5Ô+¨FÑ3Ô3ÐOÐOÐOr   r   )Údim)rM   Úget_backendr   ÚBackendÚNCCLÚXCCLrN   r   r‹   r0   r   r   rU   r=   ÚsumÚstack)rQ   ro   rO   rX   rk   Úgxss         r   rY   z_AllGather.backward›  sÐ   € åÔ #¤)Ð,Ñ,Ô,µ´Ô1BÅDÄLÔDUÐ0VÐVÐVÝ”= s¤yÐ1Ñ1Ô1ˆDÝÔ! ,¨tÔ"4Ñ5Ô5ˆBÝ Ô&¥x¤|°S´YÀÐRÀ\ÐRÐRÐRˆBˆBð PÐOÀ,ÐOÑOÔOˆKÝ”/ #¤)¨[ÐH¸<ÐHÐHÐHˆCÝ”�5œ; sÑ+Ô+°Ð3Ñ3Ô3ˆBØ�bˆzÐr   NrZ   r_   r   r   r5   r5   Œ  sH   € € € € € Øð
&ð 
&ñ „\ð
&ð ðð ñ „\ðð ð r   r5   c                   ó:   — e Zd Zed„ ¦   «         Zed„ ¦   «         ZdS )r8   c                 óf   — || _         t          j        ||                     ¦   «         |¬¦  «         |S rK   )r   rM   r7   ri   )rQ   r9   r:   r   s       r   rR   z_AllGatherBase.forward«  s5   € ð ˆŒ	ÝÔ˜m¨\×-DÒ-DÑ-FÔ-FÈeÐTÑTÔTÐTØÐr   c                 ó<  — t          j        | j        ¬¦  «        t           j        j        t           j        j        fv rÍt          j        | j        ¬¦  «        }t          |                     ¦   «         ¦  «        }|d         |z  dk    rt          d|› d|› �¦  «        ‚|d         t          j        | j        ¬¦  «        z  |d<   t          j        ||j        |j        ¬¦  «        }t          j        ||t          j        | j        ¦  «         nt          d¦  «        ‚d |d fS )NrL   r   zTensor with dimensions: z8 does not have first dimension divisible by world_size: ©ÚdeviceÚdtypezBackend not supported!)rM   r‘   r   r’   r“   r”   rh   ry   rt   r   r   Úemptyrœ   r�   Ú_reduce_scatter_baser   rU   )rQ   rW   Ú
world_sizeÚout_sizerX   s        r   rY   z_AllGatherBase.backward²  s(  € õ Ô #¤)Ð,Ñ,Ô,µ´Ô1BÅDÄLÔDUÐ0VÐVÐVÝÔ,°3´9Ð=Ñ=Ô=ˆJÝ˜K×,Ò,Ñ.Ô.Ñ/Ô/ˆHØ˜Œ{˜ZÑ'¨1Ò,Ð,Ý"ðV¨xð Vð VØISðVð Vñô ð ð # 1œ+­Ô)<À3Ä9Ð)MÑ)MÔ)MÑMˆH�Q‰KÝ”Ø Ô!3¸;Ô;Lðñ ô ˆBõ Ô% b¨+µx´|ÀSÄYÑOÔOÐOÐOåÐ7Ñ8Ô8Ð8Ø�b˜$ÐÐr   NrZ   r_   r   r   r8   r8   ª  sH   € € € € € Øðð ñ „\ðð
 ð ð  ñ „\ð ð  ð  r   r8   c                   ó:   — e Zd Zed„ ¦   «         Zed„ ¦   «         ZdS )r=   c                 ó<  ‡— || _         ˆfd„t          t          j        |¬¦  «        ¦  «        D ¦   «         | _        t          j        |¬¦  «        }t          d„ ‰D ¦   «         ¦  «        Št          j        |¬¦  «        t          j        j	        u r[t          t          j        |¬¦  «        ¦  «        D ]7}d }||k    rt          ‰¦  «        }t          j        ||         |||¬¦  «         Œ8n$t          j        |t          ‰¦  «        |¬¦  «         t          |¦  «        S )Nc                 óD   •— g | ]}‰|                               ¦   «         ‘ŒS r_   rs   )rd   re   r)   s     €r   rf   z%_AlltoAll.forward.<locals>.<listcomp>Ì  s4   ø€ ð &
ð &
ð &
Ø"#ˆG�AŒJ�OŠOÑÔð&
ð &
ð &
r   rL   c              3   ó>   K  — | ]}|                      ¦   «         V — Œd S rT   r„   r…   s     r   rv   z$_AlltoAll.forward.<locals>.<genexpr>Ð  s*   è è € Ð8Ð8¨1˜Ÿš™œÐ8Ð8Ð8Ð8Ð8Ð8r   )r   rg   rM   rh   Úinput_tensor_size_listrN   rj   r‘   r’   ÚGLOOry   r'   r<   )rQ   r   r�   r)   Úmy_rankre   Úto_sends      `   r   rR   z_AlltoAll.forwardÈ  s<  ø€ ð ˆŒ	ð&
ð &
ð &
ð &
Ý',­TÔ-@ÀuÐ-MÑ-MÔ-MÑ'NÔ'Nð&
ñ &
ô &
ˆÔ"õ ”- eÐ,Ñ,Ô,ˆÝÐ8Ð8°Ð8Ñ8Ô8Ñ8Ô8ˆåÔ %Ð(Ñ(Ô(­D¬LÔ,=Ð=Ð=Ý�4Ô.°UÐ;Ñ;Ô;Ñ<Ô<ð Jð J�Ø�Ø˜’<�<Ý" 7™mœm�GÝ”˜_¨QÔ/°¸!À5ÐIÑIÔIÐIÐIð	Jõ ŒOØÝ�W‘”Øðñ ô ð õ
 �_Ñ%Ô%Ð%r   c                 ó`   ‡— ˆfd„| j         D ¦   «         }dt          j        | j        |g‰¢R Ž z   S )Nc                 ój   •— g | ]/}t          j        |‰d          j        ‰d          j        ¬¦  «        ‘Œ0S )r   r›   )r   rž   rœ   r�   )rd   rt   ro   s     €r   rf   z&_AlltoAll.backward.<locals>.<listcomp>â  sQ   ø€ ð 
ð 
ð 
ð õ ŒKØ˜\¨!œ_Ô3¸<È¼?Ô;Pðñ ô ð
ð 
ð 
r   rn   )r¦   r=   r   r   )rQ   ro   rk   s    ` r   rY   z_AlltoAll.backwardà  sS   ø€ ð
ð 
ð 
ð 
ð Ô2ð	
ñ 
ô 
ˆð �iœo¨c¬i¸ÐTÀ|ÐTÐTÐTÑTÐTr   NrZ   r_   r   r   r=   r=   Ç  sM   € € € € € Øð&ð &ñ „\ð&ð, ðUð Uñ „\ðUð Uð Ur   r=   c                   ó:   — e Zd Zed„ ¦   «         Zed„ ¦   «         ZdS )rB   c                 ó”   — || _         |                     ¦   «         | _        || _        || _        t          j        |||||¬¦  «         |S )N)rD   rE   r   )r   rt   Ú
input_sizerD   rE   rM   rA   )rQ   r   r1   rD   rE   rC   s         r   rR   z_AlltoAllSingle.forwardì  sZ   € ð ˆŒ	ØŸš™œˆŒØ!2ˆÔØ 2ˆÔÝÔØØØ1Ø/Øð	
ñ 	
ô 	
ð 	
ð ˆr   c           	      óÔ   — t          j        | j        |j        |j        ¬¦  «        }dt
                               | j        || j        | j	        | 
                    ¦   «         ¦  «        fz   S )Nr›   )NNNN)r   rž   r®   rœ   r�   rB   r   r   rD   rE   ri   )rQ   rW   r    s      r   rY   z_AlltoAllSingle.backwardü  ss   € õ ”ØŒN ;Ô#5¸[Ô=Nð
ñ 
ô 
ˆð (Ý×!Ò!Ø”	ØØÔ&ØÔ%Ø×&Ò&Ñ(Ô(ñô ð+
ñ 
ð 	
r   NrZ   r_   r   r   rB   rB   ë  sH   € € € € € Øðð ñ „\ðð ð
ð 
ñ „\ð
ð 
ð 
r   rB   c                   ó:   — e Zd Zed„ ¦   «         Zed„ ¦   «         ZdS )rH   c                 ó�   — || _         || _        |                     t          j        ¬¦  «        }t          j        |||¬¦  «         |S )N)Úmemory_formatr~   )r   r-   rP   r   Úcontiguous_formatrM   rG   )rQ   r-   r   r    s       r   rR   z_AllReduce.forward  sD   € ð ˆŒ	ØˆŒØ—’­EÔ,C�ÑDÔDˆÝŒ˜ 2¨UÐ3Ñ3Ô3Ð3Øˆr   c                 óV   — dt                                | j        | j        |¦  «        fz   S rm   )rH   r   r-   r   r{   s     r   rY   z_AllReduce.backward  s(   € ð �z×/Ò/°´¸¼	À;ÑOÔOÐQÑQÐQr   NrZ   r_   r   r   rH   rH     sM   € € € € € Øðð ñ „\ðð ðRð Rñ „\ðRð Rð Rr   rH   )!r   r   Útorch.distributedÚdistributedrM   Útorch.autogradr   r   r   r   r   ÚWORLDr   r#   r'   rU   r+   r/   r4   r7   r<   rA   rG   r   r$   r(   r,   r0   r5   r8   r=   rB   rH   r_   r   r   ú<module>r¹      s—  ðà €€€à €€€Ø  Ð  Ð  Ð  Ð  Ð  Ø #Ð #Ð #Ð #Ð #Ð #ð
 .Ð -Ð -Ð -Ð -Ð -Ð -Ð -ð 6:ð ð ð ð ð ðð ð ð "'¤ð 0ð 0ð 0ð 0ð4  ¤ð -ð -ð -ð -ð"  %¤+ð 0ð 0ð 0ð 0ð, $œ<¨u¬{ð 1ð 1ð 1ð 1ð. +3¬,¸e¼kð Að Að Að Að8 #œ[ð +ð +ð +ð +ð. 9>¼ð #Dð #Dð #Dð #DðL =B¼Kð Jð Jð Jð Jð* ØØ
Œ+ð%ð %ð %ð %ðP #œ,¨e¬kð /ð /ð /ð /ð6 ð  ð  ð  ð  �ñ  ô  ð  ð,Sð Sð Sð Sð Sˆhñ Sô Sð Sð4Mð Mð Mð Mð Mˆxñ Mô Mð Mð*Yð Yð Yð Yð Yˆhñ Yô Yð Yð Mð Mð Mð Mð M�hñ Mô Mð Mð"ð ð ð ð �ñ ô ð ð< ð  ð  ð  ð  �Xñ  ô  ð  ð:!Uð !Uð !Uð !Uð !U�ñ !Uô !Uð !UðH
ð 
ð 
ð 
ð 
�hñ 
ô 
ð 
ðDRð Rð Rð Rð R�ñ Rô Rð Rð Rð Rr   