o
    Me                     @   s   d dl 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 ddl	mZmZmZmZ d dlmZ d dlma d d	lmZ dad
d ZdS )    N   )topology)ParallelMode)TensorParallelmodel_parallel_random_seed)PipelineParallelShardingParallelPipelineParallelWithInterleavePipelineLayer)core)_grad_scalar)fleetc                 C   s  t j }| dusJ d| dkr| S d}|j}|jdkrcd}|jd r&dnd}| dkr9tjj| ddddd	} |jd
 }|jd }|jd }|jd }|jd }	|jd }
tjj|||||	|
da	|j
dkrvtj| |j|j|jd}|S |j tjkrt| |j|d} | S |j tjkr|jdkrddlm}m} |j|j ksJ || |j tj| |j|j|jd} | S |j tjkrt| |j|d} | S |j tjkrt| tsJ d|  dkrt | |j|d} | S t!| |j|d} | S )a  
    Return distributed data parallel model (Only work in dygraph mode)

    Args:
        model (Layer): the user-defind model which inherits Layer.

    Returns:
        distributed data parallel model which inherits Layer.

    Examples:

        .. code-block:: python

            import paddle
            import paddle.nn as nn
            from paddle.distributed import fleet

            class LinearNet(nn.Layer):
                def __init__(self):
                    super(LinearNet, self).__init__()
                    self._linear1 = nn.Linear(10, 10)
                    self._linear2 = nn.Linear(10, 1)

                def forward(self, x):
                    return self._linear2(self._linear1(x))

            # 1. initialize fleet environment
            fleet.init(is_collective=True)

            # 2. create layer & optimizer
            layer = LinearNet()
            loss_fn = nn.MSELoss()
            adam = paddle.optimizer.Adam(
                learning_rate=0.001, parameters=layer.parameters())

            # 3. get data_parallel model using fleet
            adam = fleet.distributed_optimizer(adam)
            dp_layer = fleet.distributed_model(layer)

            # 4. run layer
            inputs = paddle.randn([10, 10], 'float32')
            outputs = dp_layer(inputs)
            labels = paddle.randn([10, 1], 'float32')
            loss = loss_fn(outputs, labels)

            print("loss:", loss.numpy())

            loss.backward()

            adam.step()
            adam.clear_grad()


    Nzmodel should not be Noner   FTZuse_pure_fp16ZO2ZO1)modelsZ
optimizerslevelZmaster_weightZ
save_dtypeinit_loss_scaling
incr_ratio
decr_ratioincr_every_n_stepsdecr_every_n_nan_or_infuse_dynamic_loss_scaling)r   r   r   r   r   r   )Zcomm_buffer_sizeZlast_comm_buffer_sizefind_unused_parameters)strategyr   )broadcast_mp_parametersbroadcast_sharding_parameterszDFor pipeline parallel, the model should an instance of PipelineLayer)"r   Z
worker_numZ_user_defined_strategyampZamp_configsupperpaddleZdecorateZ
GradScalerr   Zheter_ccl_modeZDataParallelZfuse_grad_size_in_MBZlast_comm_group_size_MBr   Z_hcgZget_parallel_moder   ZSHARDING_PARALLELr   ZDATA_PARALLELZsharding_degreeZ3paddle.distributed.fleet.utils.hybrid_parallel_utilr   r   Z get_sharding_parallel_world_sizeZTENSOR_PARALLELr   ZPIPELINE_PARALLEL
isinstancer
   Zget_num_virtual_stagesr   r	   )modelZ	fleet_envZ
amp_enabler   Z	amp_levelr   r   r   r   r   r   distributed_modelr   r    r    ND:\Projects\ConvertPro\env\Lib\site-packages\paddle/distributed/fleet/model.pyr      s   7






r   )r   osnumpynpbaser   tpZbase.topologyr   Zmeta_parallelr   r   r   r   r	   r
   Zpaddle.fluidr   Z*paddle.fluid.dygraph.varbase_patch_methodsr   Zpaddle.distributedr   r   r    r    r    r!   <module>   s   