
    pj                     ^    S SK r S SKJr  SSKJr  SSKJrJrJrJ	r	J
r
JrJrJrJrJr  SqS rg)    N)fleet   )ParallelMode)
DualPipeVParallelNoPipelineParallelPipelineLayerPipelineParallelPipelineParallelWithInterleave$PipelineParallelWithInterleaveFthenBSegmentParallelShardingParallelTensorParallelVPPFhenBInBalancedMemoryc           	         [         R                   nUR                  nU c   S5       e[        R                  R	                  5       S::  a  [        XS9n U $ UR                  (       a  UR                  S   (       d  UR                  S   (       a  SOSnUS:X  a8  [        R                  R                  U SSSSUR                  S   (       a  S	OS
S9n UR                  S   nUR                  S   nUR                  S   nUR                  S   nUR                  S   nUR                  S   n	[        R                  R                  UUUUUU	S9q
UR                  R                  5       [        R                  :X  Ga1  [        U [         5      (       d   S5       eUR"                  S   R$                  (       a  ['        XR                  US9n U $ U R)                  5       S:X  a  [+        XR                  US9n U $ UR,                  S   n
UR                  R/                  5       nU
SU-  :  a  [1        XR                  US9n U $ Xs=::  a	  SU-  :  aN  O  OKUR"                  S   R2                  (       a  [5        XR                  US9n U $ [7        XR                  US9n  U $ [9        SU
 SU S35      e[        U [         5      (       a  [        XUR                  S9n U $ UR:                  (       a7  [        R<                  " U UR>                  UR@                  URB                  S9nU$ UR                  R                  5       [        RD                  :X  a  [G        XR                  US9n U $ UR                  R                  5       [        RH                  :X  aP  [        R<                  " U UR>                  UR@                  URB                  UR                  RK                  5       S9n U $ UR                  R                  5       [        RL                  :X  a  [O        XR                  US9n U $ UR                  R                  5       [        RP                  :X  a  [S        XR                  US9n U $ )a  
Return distributed data parallel model (Only work in dygraph mode)

Args:
    model (Layer): the user-defined 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   )strategy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#   zDFor pipeline parallel, the model should an instance of PipelineLayer
pp_configsaccumulate_steps   zThe accumulate_steps(z/) should be greater than or equal to pp_degree())r   hcg)comm_buffer_sizelast_comm_buffer_sizefind_unused_parameters)r)   r*   r+   group)*r   _user_defined_strategypaddledistributedget_world_sizer   ampamp_configsdecorate
GradScaler_grad_scalar_hcgget_parallel_moder   PIPELINE_PARALLEL
isinstancer   hybrid_configsuse_dualpipevr   get_num_virtual_stagesr	   pipeline_configsget_pipe_parallel_world_sizer
   best_unbalanced_schedulerr   r   
ValueErrorheter_ccl_modeDataParallelfuse_grad_size_in_MBlast_comm_group_size_MBr+   SHARDING_PARALLELr   DATA_PARALLELget_data_parallel_groupSEGMENT_PARALLELr   TENSOR_PARALLELr   )model	fleet_envr   r   r   r   r    r!   r"   r#   r%   	pp_degreedistributed_models                Z/var/www/html/pdf-tiff/venv/lib/python3.13/site-packages/paddle/distributed/fleet/model.pyrM   rM   #   s]   f I//H888((*a/"5<|| ##O4##O4  	 	 D=JJ''"  ++O< # ( E %001DE)),7
)),7
%112FG"*"6"6%#
 $,#7#7&$
 
 zz,,/!!1$;%= - 
 ~~'')\-K-KK%// 	
R	
/ ""<0>>%e^^hOER LQ ))+q0$UNNXNEL LI  (889KL!CCEI1y=06>>H@ L{ >Y>** ++, 5~~Er Lk A~~Ej Lc !+,<+==lmvlwwxy  e]++&innEV LO &&$*$7$7%-%B%B*2*J*J+3+J+J	%! )( 002112 )>>H4 L- 002l6P6PP++%-%B%B*2*J*J+3+J+J#..@@B( L 002001 (>>H L 002//0 'unnxPL    )r.   paddle.distributedr   base.topologyr   meta_parallelr   r   r   r	   r
   r   r   r   r   r   r5   rM    rO   rN   <module>rT      s,     $ '   srO   