o
    Me.                     @   s8  d dl Z d dlZd dlZd dlmZ d dlZd dlmZ d dl	m
Z
 ddlmZmZmZ ddlmZ ddlmZ dd	lmZ dd
lmZ d dlmZ d dlm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l	m$Z$ ddl%m&Z&m'Z' d dl(Z(g Z)dd Z*dd Z+dd Z,ee+Z-ee,Z.G dd de/Z0dS )    N)
MethodType)_global_flags)compiler   )UserDefinedRoleMakerPaddleCloudRoleMakerRoleMakerBase)StrategyCompiler)DistributedStrategy)MetaOptimizerFactory)RuntimeFactory)wrap_decorator)parallel_helper)apply_build_strategy)topology)model_parallel_random_seed)_C_ops_legacy_C_ops)core)loggerset_log_levelc                 C   sv   |j j }t d s|S t| di }|r|d } |jd }d|ji}|j j}|r4|jr4t	
d d|_t| |||S )NZFLAGS_apply_pass_to_program_pipeline_optZsection_programstartup_programZuse_cudazCurrently, the fuse_all_optimizer_ops pass has conflict with fuse_all_reduce_ops pass. Disable the fuse_all_optimizer_ops pass temporarily.F)_user_defined_strategybuild_strategyZ_copyr   getattrr   _is_collectiveZfuse_all_reduce_opsZfuse_all_optimizer_opsr   warningr   )main_programr   configr   Zpipeline_optZ
pass_attrsZfuse_all_reduce r    ND:\Projects\ConvertPro\env\Lib\site-packages\paddle/distributed/fleet/fleet.pyapply_ir_passes,   s"   



r"   c                        fdd}|S )Nc                     s(   | d }|j d u rtd | i |S )Nr   z+Fleet can not find suitable runtime handler)_runtime_handle
ValueErrorargskwargsclsfuncr    r!   __impl__I   s   
z*_inited_runtime_handler_.<locals>.__impl__r    r+   r,   r    r*   r!   _inited_runtime_handler_H   s   r.   c                    r#   )Nc                     sB   | d }|j d ur|j  du rtd j  d S  | i |S )Nr   Tz:%s() function doesn't work when use non_distributed fleet.)_role_maker_is_non_distributedr   r   __name__r&   r*   r    r!   r,   U   s   
z,_is_non_distributed_check_.<locals>.__impl__r    r-   r    r*   r!   _is_non_distributed_check_T   s   r2   c                   @   sN  e Zd ZdZdd Z				dgddZd	d
 Zdd Zdd Zdd Z	dd Z
dd Zdd Zdd Zdd Zdd Zdd Zdd  Zdhd!d"Zd#d$ Zd%d& Zdhd'd(Zd)d* Zd+d, Zeedid-d.Zeedid/d0Zd1d2 Zeed3d4 Zeed5d6 Zeed7d8 Zeed9d: Z eed;d< Z!eed=d> Z"eed?d@ Z#eeg g fdAdBZ$ee		C	DdjdEdFZ%eedkdGdHZ&eedIdJ Z'eedKdL Z(eedMdN Z)ee	didOdPZ*didQdRZ+didSdTZ,dUdV Z-dWdX Z.	dldYdZZ/d[d\ Z0d]d^ Z1d_d` Z2	dmdadbZ3	dmdcddZ4			dmdedfZ5dS )nFleeta  
    Unified API for distributed training of PaddlePaddle
    Please reference the https://github.com/PaddlePaddle/PaddleFleetX for details


    Returns:
        Fleet: A Fleet instance

    Example for collective training:

        .. code-block:: python

            import paddle
            paddle.enable_static()
            import paddle.distributed.fleet as fleet

            fleet.init(is_collective=True)

            strategy = fleet.DistributedStrategy()
            optimizer = paddle.optimizer.SGD(learning_rate=0.001)
            optimizer = fleet.distributed_optimizer(optimizer, strategy=strategy)

            # do distributed training


    Example for parameter server training:

        .. code-block:: python

            import paddle
            paddle.enable_static()
            import paddle.distributed.fleet as fleet
            strategy = fleet.DistributedStrategy()
            fleet.init(strategy=strategy)

            optimizer = paddle.optimizer.SGD(learning_rate=0.001)
            optimizer = fleet.distributed_optimizer(optimizer)

            if fleet.is_first_worker():
                print("this is first worker")

            print("current node index: {}".format(fleet.worker_index()))
            print("total number of worker num: {}".format(fleet.worker_num()))

            if fleet.is_worker():
                print("this is worker")
            print("worker endpoints: {}".format(fleet.worker_endpoints(to_string=True)))

            print("server num: {}".format(fleet.server_num()))
            print("server endpoints: {}".format(fleet.server_endpoints(to_string=True)))

            if fleet.is_server():
                print("this is server")
            fleet.stop_worker()


    c                 C   s6   d | _ d | _d| _d | _d | _i | _tjd| _	d S )NFg        )
r/   strategy_compilerr   r$   Z_util_contextpaddle	optimizerZ	Optimizeruser_defined_optimizerselfr    r    r!   __init__   s   zFleet.__init__NFINFOc                    s  t | |du rt }t|| _|du r.t|tr%|| _t| jd| _	nt
dt|t|tr;|| _	|j| _n	t
dt|| j	  ddlm  m} |j| j	 t | _| j	 ry| jrytjj rytjj }|dkryt
dtjj r|  dkrt  | _!t"| j!| _#dS t$% rt&'d nd	t(j)v rt&'d
 n	t*| jj+t(j)d	< tj,  | jj-dkrtj.du r| /  | S t&'d | S | jro| jj0}| 1 }|  }	|rdnd}
t2t3|	}tj.du rt4  tj.}|| _#|5d||	|
| | jj6}|p
|}|du rdS d}d}|r#| jj7}t8|d }|r0| jj9}t8|d }|r=|r=||ks=J |rB|n|  dkro|	  dksRJ d}|  }|   fdd|D }|5d| || | S )aV  
        Initialize role_maker in Fleet.

        This function is responsible for the distributed architecture
        what you want to run your code behind.

        Args:
            role_maker (RoleMakerBase, optional): A ``RoleMakerBase`` containing the configuration
                of environment variables related to distributed training.If you did not initialize
                the rolemaker by yourself, it will be automatically initialized to PaddleRoleMaker.
                The default value is None.
            is_collective (Boolean, optional): A ``Boolean`` variable determines whether the program
                runs on the CPU or GPU. False means set distributed training using CPU, and True means
                GPU.The default value is False.The default value is False.
            strategy (DistributedStrategy): Extra properties for distributed training.
                For details, please refer to paddle.distributed.fleet.DistributedStrategy. Default: None.
            log_level (Integer, String, optional): A ``Integer`` or ``String`` Variable determining how hight
                the logging level is. Default is "INFO".


        Returns:
            None

        Examples1:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()

        Examples2:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init(is_collective=True)

        Examples3:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                role = fleet.PaddleCloudRoleMaker()
                fleet.init(role)

        Examples4:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                strategy = fleet.DistributedStrategy()
                fleet.init(strategy=strategy)

        Examples5:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                strategy = fleet.DistributedStrategy()
                fleet.init(log_level = "DEBUG")

        N)is_collectivez8`is_collective` should be instance of `bool`, but got {}z>`role_maker` should be subclass of `RoleMakerBase`, but got {}r   r   z[CUDA_VISIBLE_DEVICES shoule be set only 1 card if you use `python` to launch fleet program.z6The dygraph parallel environment has been initialized.ZFLAGS_nccl_nringszYou have set the environment variable FLAGS_nccl_nrings outside the program, so the nccl_comm_num in DistributedStrategy will not take effect here.Fz=The dygraph hybrid parallel environment has been initialized.   global	mp_degreeZtensor_parallel_degreec                    s   g | ]
}|  kr|qS r    r    ).0idxr@   Zmp_group_idr    r!   
<listcomp>|  s
    zFleet.init.<locals>.<listcomp>model):r   r
   copydeepcopyr   
isinstanceboolr   r   r/   r%   formattyper   Z_generate_rolepaddle.distributed.fleetdistributedfleetutilZ_set_role_makerr	   r4   r0   r6   fluidr   Zis_compiled_with_cudaZget_cuda_device_count	framework_non_static_mode
worker_numtpCommunicateTopology	_topologyHybridCommunicateGroup_hcgr   Z_is_parallel_ctx_initializedr   r   osenvironstrZnccl_comm_numZinit_parallel_envZheter_ccl_modeZ_HYBRID_PARALLEL_GROUP_init_hybrid_parallel_envshardingworker_indexlistrangeZ_CommunicateGroupZset_comm_groupZtensor_parallelsharding_configsinttensor_parallel_configs)r:   
role_makerr=   strategy	log_levelrN   Zgpus_numZuse_shardingZglobal_rankZglobal_world_sizeZglobal_ring_idZglobal_ranksZcgZuse_tensor_parallelZuse_mpZmp_degree_shardingZmp_degree_tensor_parallelra   rc   Z
mp_ring_idZmp_rankZmp_group_ranksr    rC   r!   init   s   F









GE




z
Fleet.initc                 C   s.  | j j| _| jd | _| jd | _| jd | _| jd | _| jdks&J d| jdks/J d| jdks8J dt| jd	| _t| jd	| _| jdk rYtj	 }|| j| j  | _t| jd	| _t
jg d
| j| j| j| jgd| _t
| j| _| jd	kr| j j}|d }|dkrt  dS t| dS dS )z!initialize the hybrid environment	dp_degreer@   	pp_degreesharding_degreer   z)mp_degree should be greater or equal to 0z)pp_degree should be greater or equal to 0z/sharding_degree should be greater or equal to 0r   )datapiper]   rE   )Zhybrid_group_namesdimstensor_init_seedN)r   Zhybrid_configsrh   r@   ri   rj   maxr6   rM   Zget_world_sizerT   rU   rV   rW   rX   rc   r   )r:   Znranksrc   rn   r    r    r!   r\     sB   





zFleet._init_hybrid_parallel_envc                 C      | j d usJ | j S N)rX   r9   r    r    r!   get_hybrid_communicate_group     z"Fleet.get_hybrid_communicate_groupc                 C   rq   rr   )rV   r9   r    r    r!   get_hybrid_parallel_topology  rt   z"Fleet.get_hybrid_parallel_topologyc                 C   
   | j  S )an  
        Check whether the node is the first instance of worker.

        Returns:
            bool: True if this is the first node of worker,
                  False if not.

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()
                fleet.is_first_worker()

        )r/   Z_is_first_workerr9   r    r    r!   is_first_worker     
zFleet.is_first_workerc                 C   rv   )a
  
        Get current worker index.

        Returns:
            int: node id

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()
                fleet.worker_index()

        )r/   Z_worker_indexr9   r    r    r!   r^        
zFleet.worker_indexc                 C   rv   )a  
        Get current total worker number.

        Returns:
            int: worker numbers

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()
                fleet.worker_num()

        )r/   Z_worker_numr9   r    r    r!   rS     ry   zFleet.worker_numc                 C   rv   rr   )r/   Z_get_node_numr9   r    r    r!   node_num     
zFleet.node_numc                 C   rv   rr   )r/   Z_get_local_rankr9   r    r    r!   
local_rank  r{   zFleet.local_rankc                 C   rv   rr   )r/   Z_get_local_device_idsr9   r    r    r!   local_device_ids  r{   zFleet.local_device_idsc                 C   rv   rr   )r/   Z_get_world_device_idsr9   r    r    r!   world_device_ids  r{   zFleet.world_device_idsc                 C   rv   )aY  
        Check whether the node is an instance of worker.

        Returns:
            bool: True if this is a node of worker,
                  False if not.

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()
                fleet.is_worker()

        )r/   Z
_is_workerr9   r    r    r!   	is_worker  rx   zFleet.is_workerc                 C   rv   rr   )r/   Z_is_coordinatorr9   r    r    r!   is_coordinator  r{   zFleet.is_coordinatorc                 C      |r
d | j S | j S )aQ  
        Get current worker endpoints, such as ["127.0.0.1:1001", "127.0.0.1:1002"].

        Returns:
            list/string: server endpoints

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()
                fleet.worker_endpoints()

        ,)joinr/   Z_get_trainer_endpointsr:   Z	to_stringr    r    r!   worker_endpoints  s   
zFleet.worker_endpointsc                 C   s   t | j S )a  
        Get current total worker number.

        Returns:
            int: server number

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()
                fleet.server_num()
        )lenr/   _get_pserver_endpointsr9   r    r    r!   
server_num)  s   zFleet.server_numc                 C   rv   )a
  
        Get current server index.

        Returns:
            int: node id

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()
                fleet.server_index()

        )r/   Z_server_indexr9   r    r    r!   server_index:  ry   zFleet.server_indexc                 C   r   )aQ  
        Get current server endpoints, such as ["127.0.0.1:1001", "127.0.0.1:1002"].

        Returns:
            list/string: server endpoints

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()
                fleet.server_endpoints()

        r   )r   r/   r   r   r    r    r!   server_endpointsL  s   
zFleet.server_endpointsc                 C   rv   )aY  
        Check whether the node is an instance of server.

        Returns:
            bool: True if this is a node of server,
                  False if not.

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()
                fleet.is_server()

        )r/   Z
_is_serverr9   r    r    r!   	is_serverb  rx   zFleet.is_serverc                 C   s   | j d dS )zH
        barrier all workers

        Returns:
            None
        ZworkerN)r/   Z_barrierr9   r    r    r!   barrier_workeru  s   zFleet.barrier_workerc                 C      | j | dS )ar  
        initialize `Communicator` for parameter server training.


        Returns:
            None

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()

                # build net
                # fleet.distributed_optimizer(...)

                fleet.init_worker()

        N)r$   Z_init_workerr:   Zscopesr    r    r!   init_worker~  s   zFleet.init_workerc                 C   r   )z-
        initialize coordinator node
        N)r$   Z_init_coordinatorr   r    r    r!   init_coordinator  s   zFleet.init_coordinatorc                 C   s   | j   d S rr   )r$   Z_make_fl_strategyr9   r    r    r!   make_fl_strategy  s   zFleet.make_fl_strategyc                 C   s   | j jS )z/
        get worker(training node) ptr
        )r$   Z_workerr9   r    r    r!   get_fl_client  s   zFleet.get_fl_clientc                 O   s   | j j|i | dS )a  
        init_server executor to initialize startup program,
        if the `args` is not empty, it will run load_persistables for increment training.


        Returns:
            None

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()

                # build net
                # fleet.distributed_optimizer(...)

                fleet.init_server()

        N)r$   Z_init_server)r:   r'   r(   r    r    r!   init_server  s   zFleet.init_serverc                 C      | j || dS )aa  
        load fleet model from path


        Returns:
            None

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()

                # build net
                # fleet.distributed_optimizer(...)

                fleet.load_model("path", mode=0)

        N)r$   Z_load_persistablesr:   pathmoder    r    r!   
load_model     zFleet.load_modelc                 C      | j ||| dS )al  
        load fleet one table from path


        Returns:
            None

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()

                # build net
                # fleet.distributed_optimizer(...)

                fleet.load_one_table(0, "path", mode=0)

        N)r$   Z_load_one_tabler:   Ztable_idr   r   r    r    r!   load_one_table     zFleet.load_one_tablec                 C   r   )au  
        load fleet inference model from path


        Returns:
            None

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()

                # build net
                # fleet.distributed_optimizer(...)

                fleet.load_inference_model("path", mode=1)

        N)r$   Z_load_inference_modelr   r    r    r!   load_inference_model  r   zFleet.load_inference_modelc                 C      | j   dS )a  
        run server will run pserver main program with executor.

        Returns:
            None

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()

                # build net
                # fleet.distributed_optimizer(...)

                if fleet.is_server():
                    fleet.init_server()

        N)r$   Z_run_serverr9   r    r    r!   
run_server  s   zFleet.run_serverc                 C   r   )a  
        stop `Communicator` and give training complete notice to parameter server.

        Returns:
            None

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()

                # build net
                # fleet.distributed_optimizer(...)

                fleet.init_server()

        N)r$   Z_stop_workerr9   r    r    r!   stop_worker(  s   zFleet.stop_workerc              	   K   s  d}|s|sd}t  }t j|}|rog }g }	|D ]}
t|
tr'||
 qt|
t jjr5||
j qt	d|D ]}
t|
trH|	|
 q;t|
t jjrV|	|
j q;t	ddd |	D }| j
||||d dd d S d}d|v r{t|d }| j
j||d |d d S )	NTFzfeed must be [str|Variable]c                 S   s    g | ]}t j  |qS r    )r6   staticdefault_main_programZglobal_blockvar)rA   namer    r    r!   rD   _  s    zFleet.save.<locals>.<listcomp>r   r   )r   r   )r6   ZCPUPlacer   ExecutorrH   r[   appendVariabler   r%   r$   _save_inference_modelrb   _save_persistables)r:   dirnamefeedfetchconfigsZ	inferenceplaceexecutorfeeded_var_namesZfetch_var_namesr   Z
fetch_varsZincrement_moder    r    r!   save@  s@   


z
Fleet.saveTr   c              	   C   s   | j ||||||| dS )a\  
        save inference model for inference.

        Returns:
            None

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()

                # build net
                # fleet.distributed_optimizer(...)

                fleet.init_server()

        N)r$   r   )r:   r   r   r   Ztarget_varsr   Zexport_for_deploymentr   r    r    r!   save_inference_modelo  s    zFleet.save_inference_modelc                 C   s   | j |||| dS )aX  

        saves all persistable tensors from :code:`main_program` to
        the folder :code:`dirname`. You can refer to

        The :code:`dirname` is used to specify the folder where persistable tensors
        are going to be saved. If you would like to save tensors in separate
        files, set :code:`filename` None.

        Args:
            executor(Executor): The executor to run for saving persistable tensors.
                                You can refer to :ref:`api_guide_executor_en` for
                                more details.

            dirname(str, optional): The saving directory path.
                                When you need to save the parameter to the memory, set it to None.
            main_program(Program, optional): The program whose persistbale tensors will
                                             be saved. Default: None.


        Returns:
            None

        Examples:

            .. code-block:: text

                import paddle
                paddle.enable_static()
                import paddle.distributed.fleet as fleet

                fleet.init()

                # build net
                # fleet.distributed_optimizer(...)

                exe = paddle.static.Executor(paddle.CPUPlace())
                fleet.save_persistables(exe, "dirname", paddle.static.default_main_program())

        N)r$   r   )r:   r   r   r   r   r    r    r!   save_persistables  s   +zFleet.save_persistablesc                 K   s   | j j|fi |S rr   )r$   Z_save_cache_model)r:   r   r   r    r    r!   save_cache_model  s   zFleet.save_cache_modelc                 C   rv   rr   )r$   Z_check_save_pre_patch_doner9   r    r    r!   check_save_pre_patch_done  s   
zFleet.check_save_pre_patch_donec                 C   r   )al  
        save fleet one table from path


        Returns:
            None

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()

                # build net
                # fleet.distributed_optimizer(...)

                fleet.save_one_table(0, "path", mode=0)

        N)r$   Z_save_one_tabler   r    r    r!   save_one_table  r   zFleet.save_one_tablec                 C   s   | j ||||| dS )a<  
        save fleet one table from path


        Returns:
            None

        Examples:

            .. code-block:: python

                import paddle.distributed.fleet as fleet
                fleet.init()
                import paddle
                place = paddle.fluid.CPUPlace()
                exe = paddle.fluid.Executor(place)

                # build net
                # fleet.distributed_optimizer(...)

                fleet.save_dense_params(exe, "path", scope=paddle.static.global_scope(), program=paddle.static.default_main_program())

        N)r$   Z_save_dense_params)r:   r   r   scopeprogramZ	var_namesr    r    r!   save_dense_params  s   
zFleet.save_dense_paramsc                 C   s   | j | d S rr   )r$   Z_shrink)r:   	thresholdr    r    r!   shrink  s   zFleet.shrinkc                 C   s4   || _ |dur| jrtd t|| _i | _| S )a  
        Optimizer for distributed training.

        For the distributed training, this method would rebuild a new instance of DistributedOptimizer.
        Which has basic Optimizer function and special features for distributed training.

        Args:
            optimizer(Optimizer): The executor to run for init server.
            strategy(DistributedStrategy): Extra properties for distributed optimizer.
                It is recommended to use DistributedStrategy in fleet.init(). The strategy
                here is for compatibility. If the strategy in fleet.distributed_optimizer()
                is not None, then it will overwrite the DistributedStrategy in fleet.init(),
                which will take effect in distributed training.

        Returns:
            Fleet: instance of fleet.

        Examples:

            .. code-block:: python

                import paddle
                import paddle.distributed.fleet as fleet
                fleet.init(is_collective=True)
                strategy = fleet.DistributedStrategy()
                optimizer = paddle.optimizer.SGD(learning_rate=0.001)
                optimizer = fleet.distributed_optimizer(optimizer, strategy=strategy)

        Na  It is recommended to use DistributedStrategy in fleet.init(). The strategy here is only for compatibility. If the strategy in fleet.distributed_optimizer() is not None, then it will overwrite the DistributedStrategy in fleet.init(), which will take effect in distributed training.)r8   r   r   r   rF   rG   r   r5   )r:   r7   re   r    r    r!   distributed_optimizer  s   zFleet.distributed_optimizerc                 C   sT   d }| j  D ]}t|dr|} nq|d u r t| jdr | j}|d us(J d|S )Namp_initzSamp_init can only be used when the amp(auto mixed precision) strategy is turned on.)r4   Z_get_applied_meta_optimizerhasattrr8   )r:   amp_optimizerr7   r    r    r!   _get_amp_optimizer=  s   

zFleet._get_amp_optimizerc                 C   s   |   }| S )z)Return the real-time loss scaling factor.)r   get_loss_scaling)r:   r   r    r    r!   r   N  s   zFleet.get_loss_scalingc                 C   s   |   }|||||S )a  
        Init the amp training, such as cast fp32 parameters to fp16 type.

        Args:
            place(CUDAPlace): place is used to initialize
                fp16 parameters with fp32 values.
            scope(Scope): The scope is used to find fp32 parameters.
            test_program(Program): The program is used for testing.
            use_fp16_test(bool): Whether to use fp16 testing.

        Examples:
            .. code-block:: python

                import paddle
                import paddle.nn.functional as F
                paddle.enable_static()

                def run_example_code():
                    place = paddle.CUDAPlace(0)
                    exe = paddle.static.Executor(place)
                    data = paddle.static.data(name='X', shape=[None, 1, 28, 28], dtype='float32')
                    conv2d = paddle.static.nn.conv2d(input=data, num_filters=6, filter_size=3)
                    # 1) Use fp16_guard to control the range of fp16 kernels used.
                    with paddle.static.amp.fp16_guard():
                        bn = paddle.static.nn.batch_norm(input=conv2d, act="relu")
                        pool = F.max_pool2d(bn, kernel_size=2, stride=2)
                        hidden = paddle.static.nn.fc(pool, size=10)
                        loss = paddle.mean(hidden)
                    # 2) Create the optimizer and set `multi_precision` to True.
                    # Setting `multi_precision` to True can avoid the poor accuracy
                    # or the slow convergence in a way.
                    optimizer = paddle.optimizer.Momentum(learning_rate=0.01, multi_precision=True)
                    # 3) These ops in `custom_black_list` will keep in the float32 computation type.
                    amp_list = paddle.static.amp.CustomOpLists(
                        custom_black_list=['pool2d'])
                    # 4) The entry of Paddle AMP.
                    # Enable pure fp16 training by setting `use_pure_fp16` to True.
                    optimizer = paddle.static.amp.decorate(
                        optimizer,
                        amp_list,
                        init_loss_scaling=128.0,
                        use_dynamic_loss_scaling=True,
                        use_pure_fp16=True)
                    # If you don't use the default_startup_program(), you sholud pass
                    # your defined `startup_program` into `minimize`.
                    optimizer.minimize(loss)
                    exe.run(paddle.static.default_startup_program())
                    # 5) Use `amp_init` after FP32 parameters initialization(such as `exe.run(startup_program)`).
                    # If you want to perform the testing process, you should pass `test_program` into `amp_init`.
                    optimizer.amp_init(place, scope=paddle.static.global_scope())

                if paddle.is_compiled_with_cuda() and len(paddle.static.cuda_places()) > 0:
                    run_example_code()
        )r   r   )r:   r   r   Ztest_programZuse_fp16_testr   r    r    r!   r   S  s   9zFleet.amp_initc                 C   s    d| j vrtd i S | j d S )Nvalid_strategyzNWARNING: You may need to call minimize function before this function is calledr5   printr9   r    r    r!   _final_strategy     

zFleet._final_strategyc                 C       d| j vrtd g S | j d S )Napplied_meta_listzTWARNING: You may need to call minimize function before _get_applied_meta_list calledr   r9   r    r    r!   _get_applied_meta_list  r   zFleet._get_applied_meta_listc                 C   r   )Napplied_graph_listzUWARNING: You may need to call minimize function before _get_applied_graph_list calledr   r9   r    r    r!   _get_applied_graph_list  r   zFleet._get_applied_graph_listc                 C   sN   t |ts| ||||S tjj s| j s| j	rt
d| ||||S )a
  
        Add distributed operations to minimize ``loss`` by updating ``parameter_list``.

        Args:
            loss (Tensor): A ``Tensor`` containing the value to minimize.
            startup_program (Program, optional): :ref:`api_fluid_Program` for
                initializing parameters in ``parameter_list``. The default value
                is None, at this time :ref:`api_fluid_default_startup_program` will be used.
            parameter_list (Iterable, optional): Iterable of ``Tensor`` or ``Tensor.name`` to update
                to minimize ``loss``. The default value is None, at this time all parameters
                will be updated.
            no_grad_set (set, optional): Set of ``Tensor``  or ``Tensor.name`` that don't need
                to be updated. The default value is None.

        Returns:
            tuple: tuple (optimize_ops, params_grads), A list of operators appended
            by minimize and a list of (param, grad) tensor pairs, param is
            ``Parameter``, grad is the gradient value corresponding to the parameter.
            The returned tuple can be passed to ``fetch_list`` in ``Executor.run()`` to
            indicate program pruning. If so, the program will be pruned by ``feed`` and
            ``fetch_list`` before run, see details in ``Executor``.

        Examples:

            .. code-block:: python

                import paddle
                paddle.enable_static()
                import paddle.distributed.fleet as fleet
                import paddle.nn.functional as F

                hid_dim = 10
                label_dim = 2
                input_x = paddle.static.data(name='x', shape=[None, 13], dtype='float32')
                input_y = paddle.static.data(name='y', shape=[None, 1], dtype='int64')
                fc_1 = paddle.static.nn.fc(x=input_x, size=hid_dim, activation='tanh')
                fc_2 = paddle.static.nn.fc(x=fc_1, size=hid_dim, activation='tanh')
                prediction = paddle.static.nn.fc(x=[fc_2], size=label_dim, activation='softmax')
                cost = F.cross_entropy(input=prediction, label=input_y)
                avg_cost = paddle.mean(x=cost)

                fleet.init(is_collective=True)
                strategy = fleet.DistributedStrategy()
                optimizer = paddle.optimizer.SGD(learning_rate=0.001)
                optimizer = fleet.distributed_optimizer(optimizer, strategy=strategy)
                optimizer.minimize(avg_cost)

                # for more examples, please reference https://github.com/PaddlePaddle/PaddleFleetX

        z loss can be list only in PS mode)rH   r_   _minimize_implr6   rP   rQ   rR   r/   r0   r   r%   _minimize_losses_impl)r:   lossr   parameter_listno_grad_setr    r    r!   minimize  s   
5
zFleet.minimizec                 C   s  i }t | j|d< tjj r| j}|| _|	|S |j
j| _t| jdsVt| jdt  | jjd | jjd< | jjd | jjd< | jjd | jjd< | jjd | jjd< | j|d< | jg|d< ||d	< |d krytj jd
d| _tj }n|jd
d| _||d< |g|d< | j|d< | jjs| jjrddlm} || }|||||\}	}
}}|	|
||fS t | j}t | j|d< t | j}| r|D ]}||| qg }g }g }|D ]/}| || j| j| |! r|" s|#| q|! r|" r|#| q|#| q| j$%|| j| j|||\}}| j$&||}t ||d< t'(dt)|d   t'(dt)|d   | j$* }| j$+ }||d< ||d< || _|| _,| j,-  g }	g }
| j. r| j/s| j0d u ryt1 2|| _0t34| jj5|j6d d}||j
j_7| jj	||||dS |rt'(dt)t8|j
j  |j	||||d\}	}
t'(dt)t8|j
j  tj9 }t'(dt)t8|  t8|t8|j
jkrtjj:|j
j t'(dt)t8|  n| jj	||||d\}	}
|	|d< |
|d< |r$t'(dt)t8|j
j  |j	||||d\}	}
|	|d< |
|d < nt;|j
j||  | jj<shtj9 }|j=d u r>i n|j=}| > |d!< | ? |d"< | jj@A D ]\}}|s_||vrc|||< qS||_=| j0d u rut1 2|| _0d#d lBmC  mD} |jEF|d  |	|
fS )$Nuser_defined_strategydistributed_info_rh   r@   ri   rj   origin_main_programorigin_main_programsr   FZfor_testorigin_startup_programorigin_startup_programsrd      )AutoParallelizerr   zvalid_strategy: zuser_defined_strategy: r   r   )Z	loss_nameZshare_vars_fromr   zbefore minimize program id: zafter minimize program id: zdefault program id: z!default program id after switch: program_optimize_opsprogram_params_gradsz"before graph minimize program id: Zgraph_optimize_opsZgraph_optimize_gradsmpi_sizempi_rankr   )GrF   rG   r   r6   rP   rQ   rR   r8   r5   r   blockr   r   r   setattrdictra   r   r   default_startup_programcloner   r/   Z	semi_autoZauto_searchZauto_parallel.parallelizerr   Zparallelizer   Z_get_valid_meta_optimizersZ_is_strict_autoZ_enable_strategy_set_basic_infoZ
_can_applyZ_is_graph_outr   r4   Zgenerate_optimizerZ_get_valid_strategyr   debugr[   r   r   r   _enable_envr0   r   r$   r   _create_runtimer   ZCompiledProgramZwith_data_parallelr   Z_graphidr   Zswitch_main_programr"   Z_is_heter_parameter_server_mode
_fleet_optrS   r^   trainer_desc_configsitemsrL   rM   rN   rO   _set_strategy)r:   r   r   r   r   contextZ
target_optr   Zauto_parallelizeroptimize_opsparams_gradsZdist_startup_progZdist_main_progZdistributed_optimizer_listZcopy_user_defined_strategyoptZvalid_optimizer_listZvalid_graph_optimizer_listZcan_not_apply_optimizer_listZmeta_optimizerZgraph_optimizerr   r   r   Zcompiled_programZdefault_programr   opt_infokvrN   r    r    r!   r     sV  





















zFleet._minimize_implc                 C   s  i }|d j j| _| j|d< g |d< |D ]}|d |j j q||d< |d u r9t|dkr5tj g}ntd|d j	dd| _
|d |d	< g |d
< |D ]	}|d
 | qN| j|d< t| j|d< t| j|d< || _|d | _| j  g }g }	ddlm}
 |
| j}||| j| j| j |j||||d\}}	||d< |	|d< |D ]D}|j j}|jd u ri n|j}|  |d< |  |d< | jj D ]\}}|s||vr|||< q||_tdtt| t|j  q| j d u rt! "|| _ dd l#m$  m%} |j&'|d  ||	fS )Nr   r   r   r   r   z0startup_program can't be None when loss is list.Fr   r   r   rd   r   r   )ParameterServerOptimizerr   r   r   r   r   zfleet base opt info: )(r   r   r   r   r   r6   r   r   r%   r   r   r/   rF   rG   r   r5   r   r   Zmeta_optimizersr   r8   r   Zminimize_losses_implr   rS   r^   r   r   r   r   r[   r   r$   r   r   rL   rM   rN   rO   r   )r:   ZlossesZstartup_programsr   r   r   r   r   r   r   r   Zps_optimizerr   r   r   rN   r    r    r!   r     s   





	


zFleet._minimize_losses_impl)NFNr<   )Frr   )NTr   )Nr   )NNF)NNN)6r1   
__module____qualname____doc__r;   rg   r\   rs   ru   rw   r^   rS   rz   r|   r}   r~   r   r   r   r   r   r   r   r   is_non_distributed_checkinited_runtime_handlerr   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r    r    r    r!   r3   k   s    :
 X-

	-(-

/
<		

E
 _r3   )1rF   r6   rY   typesr   numpynpZpaddle.fluid.frameworkr   Zpaddle.fluidr   Zbase.role_makerr   r   r   Zbase.strategy_compilerr	   Zbase.distributed_strategyr
   Zbase.meta_optimizer_factoryr   Zbase.runtime_factoryr   Zpaddle.fluid.wrapped_decoratorr   Zpaddle.fluid.dygraphr   Zpaddle.fluid.irr   baser   rT   Zmeta_parallelr   r   r   r   Zutils.log_utilr   r   logging__all__r"   r.   r2   r   r   objectr3   r    r    r    r!   <module>   s8   