Ë
    ‡\;jU  ã                   óZ   — d dl Zd dlZd dlmZ d dlmZ d dlmZ ddl	m
Z
mZ dd„Z	 d	d„Zy)
é    N)Ú	framework)Ústreamé   )Úconvert_object_to_tensorÚconvert_tensor_to_objectc                 ó4   — t        j                  | ||||«      S )ak  

    Scatter a tensor to all participators. As shown below, one process is started with a GPU and the source of the scatter
    is GPU0. Through scatter operator, the data in GPU0 will be sent to all GPUs averagely.

    .. image:: https://githubraw.cdn.bcebos.com/PaddlePaddle/docs/develop/docs/api/paddle/distributed/img/scatter.png
        :width: 800
        :alt: scatter
        :align: center

    Args:
        tensor (Tensor): The output Tensor. Its data type
            should be float16, float32, float64, int32, int64, int8, uint8, bool or bfloat16.
        tensor_list (list|tuple): A list/tuple of Tensors to scatter. Every element in the list must be a Tensor whose data type
            should be float16, float32, float64, int32, int64, int8, uint8, bool or bfloat16. Default value is None.
        src (int): The source rank id. Default value is 0.
        group (Group, optional): The group instance return by new_group or None for global default group.
        sync_op (bool, optional): Whether this op is a sync op. The default value is True.

    Returns:
        None.

    Examples:
        .. code-block:: python

            >>> # doctest: +REQUIRES(env: DISTRIBUTED)
            >>> import paddle
            >>> import paddle.distributed as dist

            >>> dist.init_parallel_env()
            >>> if dist.get_rank() == 0:
            ...     data1 = paddle.to_tensor([7, 8, 9])
            ...     data2 = paddle.to_tensor([10, 11, 12])
            ...     dist.scatter(data1, src=1)
            >>> else:
            ...     data1 = paddle.to_tensor([1, 2, 3])
            ...     data2 = paddle.to_tensor([4, 5, 6])
            ...     dist.scatter(data1, tensor_list=[data1, data2], src=1)
            >>> print(data1, data2)
            >>> # [1, 2, 3] [10, 11, 12] (2 GPUs, out for rank 0)
            >>> # [4, 5, 6] [4, 5, 6] (2 GPUs, out for rank 1)
    )r   Úscatter)ÚtensorÚtensor_listÚsrcÚgroupÚsync_ops        úqG:\00. PROJECTS\API\Inventory\templateJSON\kerjaOCR\Lib\site-packages\paddle/distributed/communication/scatter.pyr	   r	      s   € ôV �>‰>˜& +¨s°E¸7ÓCÐCó    c                 óV  — t        j                  «       sJ d«       ‚t        j                  «       }g }g }||k(  rC|D ]2  }t	        |«      \  }}	|j                  |«       |j                  |	«       Œ4 t        |«      }
nt        j                  g d¬«      }
t        j                  |
|«       t        |
j                  «       «      }g }|D ]O  }|j                  «       }t        j                  ||g«      }t        j                   |«      }|j                  |«       ŒQ t        j                  |gd¬«      }t#        |||k(  r|nd||«       t        j                  g d¬«      }t#        |||k(  r|nd||«       | j%                  «        | j                  t'        ||j                  «       «      «       y)aô  

    Scatter picklable objects from the source to all others. Similiar to scatter(), but python object can be passed in.

    Args:
        out_object_list (list): The list of objects to store the scattered objects.
        in_object_list (list): The list of objects to scatter. Only objects on the src rank will be scattered.
        src (int): The source rank in global view.
        group (Group): The group instance return by new_group or None for global default group.

    Returns:
        None.

    Warning:
        This API only supports the dygraph mode.

    Examples:
        .. code-block:: python

            >>> # doctest: +REQUIRES(env: DISTRIBUTED)
            >>> import paddle.distributed as dist

            >>> dist.init_parallel_env()
            >>> out_object_list = []
            >>> if dist.get_rank() == 0:
            ...     in_object_list = [{'foo': [1, 2, 3]}, {'foo': [4, 5, 6]}]
            >>> else:
            ...     in_object_list = [{'bar': [1, 2, 3]}, {'bar': [4, 5, 6]}]
            >>> dist.scatter_object_list(out_object_list, in_object_list, src=1)
            >>> print(out_object_list)
            >>> # [{'bar': [1, 2, 3]}] (2 GPUs, out for rank 0)
            >>> # [{'bar': [4, 5, 6]}] (2 GPUs, out for rank 1)
    z6scatter_object_list doesn't support static graph mode.Úint64)ÚdtypeÚuint8N)r   Úin_dynamic_modeÚdistÚget_rankr   ÚappendÚmaxÚpaddleÚemptyr   Ú	broadcastÚintÚitemÚnumpyÚnpÚresizeÚ	to_tensorr	   Úclearr   )Úout_object_listÚin_object_listr   r   ÚrankÚin_obj_tensorsÚin_obj_sizesÚobjÚ
obj_tensorÚobj_sizeÚmax_obj_size_tensorÚmax_obj_sizeÚin_tensor_listr
   Ú
numpy_dataÚ	in_tensorÚ
out_tensorÚout_tensor_sizes                     r   Úscatter_object_listr3   J   s‚  € ôJ 	×!Ñ!Ô#ð@à?ó@Ø#ô �=‰=‹?€DØ€NØ€Làˆs‚{Û!ˆCÜ#;¸CÓ#@Ñ ˆJ˜Ø×!Ñ! *Ô-Ø×Ñ Õ)ð "ô " ,Ó/Ñä$Ÿl™l¨2°WÔ=ÐÜ
×ÑÐ(¨#Ô.ÜÐ*×/Ñ/Ó1Ó2€Lð €NÛ ˆØ—\‘\“^ˆ
Ü—Y‘Y˜z¨L¨>Ó:ˆ
Ü×$Ñ$ ZÓ0ˆ	Ø×Ñ˜iÕ(ð	 !ô
 —‘˜|˜n°GÔ<€JÜˆJ¨$°#ª+™¸4ÀÀeÔLä—l‘l 2¨WÔ5€OÜˆO¨T°Sª[™\¸dÀCÈÔOà×ÑÔØ×ÑÜ  ¨_×-AÑ-AÓ-CÓDõr   )Nr   NT)Nr   N)r   r    r   Úpaddle.distributedÚdistributedr   r   Ú paddle.distributed.communicationr   Úserialization_utilsr   r   r	   r3   © r   r   Ú<module>r9      s-   ðó ã Ý !Ý Ý 3÷ó+Dð^ 8<ôGr   