Ë
    ‡\;j�  ã                   óR   — d dl Z d dlmZ ddlmZ ddlmZmZmZm	Z	m
Z
mZmZ dad„ Zy)é    N)Úfleeté   )ÚParallelMode)ÚPipelineLayerÚPipelineParallelÚPipelineParallelWithInterleaveÚ$PipelineParallelWithInterleaveFthenBÚSegmentParallelÚShardingParallelÚTensorParallelc           	      ó6  — t         j                   }| €J d«       ‚t        j                  j                  «       dk  r| S |j                  }|j
                  rÜ|j                  d   s|j                  d   rdnd}|dk(  r6t        j
                  j                  | dddd|j                  d   rdnd	¬
«      } |j                  d   }|j                  d   }|j                  d   }|j                  d   }|j                  d   }|j                  d   }	t        j
                  j                  ||||||	¬«      a	|j                  r9t        j                  | |j                  |j                  |j                  ¬«      }
|
S |j                  j!                  «       t"        j$                  k(  rt'        | |j                  |¬«      } | S |j                  j!                  «       t"        j(                  k(  rRt        j                  | |j                  |j                  |j                  |j                  j+                  «       ¬«      } | S |j                  j!                  «       t"        j,                  k(  rt/        | |j                  |¬«      } | S |j                  j!                  «       t"        j0                  k(  rt3        | |j                  |¬«      } | S |j                  j!                  «       t"        j4                  k(  r¬t7        | t8        «      sJ d«       ‚| j;                  «       dk(  rt=        | |j                  |¬«      } | S |j>                  d   }|j                  jA                  «       }||k\  r"||dz  k  rtC        | |j                  |¬«      } | S tE        | |j                  |¬«      } | S )ag  
    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().__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   Úuse_pure_fp16Úuse_pure_bf16ÚO2ÚO1Úfloat16Úbfloat16)ÚmodelsÚ
optimizersÚlevelÚmaster_weightÚ
save_dtypeÚ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   )Úcomm_buffer_sizeÚlast_comm_buffer_sizeÚfind_unused_parameters)Ústrategy)r    r!   r"   ÚgroupzDFor pipeline parallel, the model should an instance of PipelineLayerÚaccumulate_stepsé   )#r   ÚpaddleÚdistributedÚget_world_sizeÚ_user_defined_strategyÚampÚamp_configsÚdecorateÚ
GradScalerÚ_grad_scalarÚheter_ccl_modeÚDataParallelÚfuse_grad_size_in_MBÚlast_comm_group_size_MBr"   Ú_hcgÚget_parallel_moder   ÚSHARDING_PARALLELr   ÚDATA_PARALLELÚget_data_parallel_groupÚSEGMENT_PARALLELr
   ÚTENSOR_PARALLELr   ÚPIPELINE_PARALLELÚ
isinstancer   Úget_num_virtual_stagesr   Úpipeline_configsÚget_pipe_parallel_world_sizer	   r   )ÚmodelÚ	fleet_envr#   r   r   r   r   r   r   r   Údistributed_modelr%   Ú	pp_degrees                úgG:\00. PROJECTS\API\Inventory\templateJSON\kerjaOCR\Lib\site-packages\paddle/distributed/fleet/model.pyrB   rB       s†  € ôf —‘€IàÐÐ8Ð8Ó8ÐÜ×Ñ×(Ñ(Ó*¨aÒ/Øˆà×/Ñ/€HØ‡|‚|ð ×#Ñ# OÒ4Ø×#Ñ# OÒ4ñ ð ð	 	ð �DŠ=Ü—J‘J×'Ñ'ØØØØ"Øà×'Ñ'¨Ò8ñ  àð (ó 	ˆEð %×0Ñ0Ð1DÑEÐØ×)Ñ)¨,Ñ7ˆ
Ø×)Ñ)¨,Ñ7ˆ
Ø%×1Ñ1Ð2FÑGÐØ"*×"6Ñ"6Ø%ñ#
Ðð $,×#7Ñ#7Ø&ñ$
Ð ô
 —z‘z×,Ñ,Ø/Ø!Ø!Ø1Ø$;Ø%=ð -ó 
ˆð ×ÒÜ"×/Ñ/ØØ%×:Ñ:Ø"*×"BÑ"BØ#+×#BÑ#Bô	
Ðð !Ð à‡~�~×'Ñ'Ó)¬\×-KÑ-KÒKÜ  ¨	¯©ÀÔJˆðL €LðK 
�‰×	)Ñ	)Ó	+¬|×/IÑ/IÒ	IÜ×#Ñ#ØØ%×:Ñ:Ø"*×"BÑ"BØ#+×#BÑ#BØ—.‘.×8Ñ8Ó:ô
ˆðH €Lð; 
�‰×	)Ñ	)Ó	+¬|×/LÑ/LÒ	LÜ  y§~¡~ÀÔIˆð8 €Lð7 
�‰×	)Ñ	)Ó	+¬|×/KÑ/KÒ	KÜ˜u i§n¡n¸xÔHˆð4 €Lð3 
�‰×	)Ñ	)Ó	+¬|×/MÑ/MÒ	MÜØ”=ô
ð 	RàQó	Rð 
ð ×'Ñ'Ó)¨QÒ.ä$ U¨I¯N©NÀXÔNˆEð& €Lð#  (×8Ñ8Ð9KÑLÐØ!Ÿ™×CÑCÓEˆIà  IÒ-Ø$ y°1¡}Ò4ô =Ø˜9Ÿ>™>°Hô�ð €Lô	 7Ø˜9Ÿ>™>°Hô�ð €Ló    )r'   Úpaddle.distributedr   Úbase.topologyr   Úmeta_parallelr   r   r   r	   r
   r   r   r/   rB   © rE   rD   Ú<module>rJ      s,   ðó Ý $å '÷÷ ñ ð €óSrE   