
      iEN                        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 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 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# d dl$m%Z%m&Z&m'Z'm(Z( d dl$m)Z* d dl+m,Z,m-Z- d dl.m/Z/ d dl0m1Z1 d dl2m3Z3 d dl4m5Z5 d dl6m7Z7m8Z8m9Z9 d dl:m;Z; d dl<m=Z=m>Z> d dl?m@Z@ d dlAmBZBmCZC d dlDmEZE d dlFmGZG d dlHmIZImJZJmZ erd dlKmLZL  e j                  eN      ZOdZP G d d e@      ZQ G d! d"eC      ZRy)#    N)nullcontext)	timedelta)TYPE_CHECKINGAnyCallableLiteralOptionalUnion)rank_zero_only)Tensor)Module)DistributedDataParallel)	Optimizer)override)CheckpointIOClusterEnvironment)default_pg_timeout)_StrategyRegistry)_distributed_is_initialized-_get_default_process_group_backend_for_device_init_dist_connection_sync_ddp_if_availablegroup)_IS_WINDOWS_TORCH_GREATER_EQUAL_2_3)_optimizers_to_device)
reset_seed)ReduceOp)LightningOptimizer)_register_ddp_comm_hook_sync_module_statesprepare_for_backward)	Precision)_MultiProcessingLauncher_SubprocessScriptLauncher)ParallelStrategy)
TBroadcast_ForwardRedirection)	TrainerFn_augment_message)rank_zero_deprecationrank_zero_infor   )ModelAverager)ddp_fork%ddp_fork_find_unused_parameters_false$ddp_fork_find_unused_parameters_trueddp_notebook)ddp_notebook_find_unused_parameters_false(ddp_notebook_find_unused_parameters_truec                       e Zd ZdZddddddddddedfded   deeej                        dee	   dee
   d	ee   d
ee   dee   dee   dee   dee   dee   ded   deddf fdZedefd       Zeedej                  fd              Zedefd       Zej4                  deddfd       Zedefd       Zeedeeef   fd              Zedee   fd       Zed?d       Zed? fd       Z ed@d       Z!ed e"de#fd!       Z$d?d"Z%defd#Z&d?d$Z'd?d%Z(d?d&Z)e	 dAd'e*d(eg ef   d ee+d)e"f      dedef
 fd*       Z,d?d+Z-deee      fd,Z.ed-ededdfd.       Z/edBd/e0d0ede0fd1       Z1ed2e2ddfd3       Z3ed2e2ddfd4       Z4ed?d5       Z5e	 dCd6e2d7ee   d8ee+e6ef      de2fd9       Z7e8ed:e9ddfd;              Z:ed<e;ddfd=       Z<ed? fd>       Z= xZ>S )DDDPStrategyzKStrategy for multi-process single-device training on one or multiple nodes.Npopenacceleratorzpl.accelerators.Acceleratorparallel_devicescluster_environmentcheckpoint_ioprecision_pluginddp_comm_stateddp_comm_hookddp_comm_wrappermodel_averaging_periodprocess_group_backendtimeoutstart_method)r8   spawnfork
forkserverkwargsreturnc                 >   t         |   |||||       t        j                  | j                  j
                   d       t               | _        d| _        || _	        || _
        || _        || _        |	| _        d | _        |
| _        || _        || _        d| _        y )N)r9   r:   r;   r<   r=   z: initializing DDP strategy   F)super__init__logdebug	__class____name___DDPForwardRedirection_forward_redirection
_num_nodes_ddp_kwargs_ddp_comm_state_ddp_comm_hook_ddp_comm_wrapper_model_averaging_period_model_averager_process_group_backend_timeout_start_method_pl_static_graph_delay_done)selfr9   r:   r;   r<   r=   r>   r?   r@   rA   rB   rC   rD   rH   rP   s                 u/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/lightning/pytorch/strategies/ddp.pyrM   zDDPStrategy.__init__G   s      	#- 3'- 	 	
 			T^^,,--HIJ$:$<!!-+!1'=$8<5J#-4)+0(    c                 L    t        dt        |       j                   dd       y)z1Legacy property kept for backwards compatibility.`z3.is_distributed` is deprecated. Use is discouraged.   )
stacklevelT)r-   typerQ   r_   s    r`   is_distributedzDDPStrategy.is_distributedl   s,     	T
##$$WXef	
 ra   c                 P    | j                   J | j                   | j                     S N)r:   
local_rankrg   s    r`   root_devicezDDPStrategy.root_devicet   s+     $$000$$T__55ra   c                     | j                   S rj   rT   rg   s    r`   	num_nodeszDDPStrategy.num_nodesz   s    ra   ro   c                     || _         y rj   rn   )r_   ro   s     r`   ro   zDDPStrategy.num_nodes~   s     $ra   c                 H    | j                   t        | j                         S dS Nr   )r:   lenrg   s    r`   num_processeszDDPStrategy.num_processes   s$    -1-B-B-Ns4(()UTUUra   c                 N    | j                   | j                  z  | j                  dS )N)num_replicasrank)ro   rt   global_rankrg   s    r`   distributed_sampler_kwargsz&DDPStrategy.distributed_sampler_kwargs   s$     "&$2D2D!DtO_O_``ra   c                     | j                   S rj   )r[   rg   s    r`   rB   z!DDPStrategy.process_group_backend   s    ***ra   c                     | j                   J | j                  dk(  r1t        | j                   | j                  | j                        | _        y t        | | j                        | _        y )Nr8   )rD   )r;   r]   r&   rt   ro   	_launcherr%   rg   s    r`   _configure_launcherzDDPStrategy._configure_launcher   sZ    ''333(6t7O7OQUQcQceiesestDN5dI[I[\DNra   c                 B    t         |           | j                          y rj   )rL   setup_environmentsetup_distributed)r_   rP   s    r`   r   zDDPStrategy.setup_environment   s    !# ra   c                 6   | j                   J | j                   j                  |       |j                  j                  }| j                  J |t
        j                  k(  r6| j                  r*| j                  j                  | j                        | _        | j                  j                  | j                         | j                          |t
        j                  k(  r"| j                          | j                  |       nt        | j                         | j                          |t
        j                  k(  rat!        | j"                  | j$                         dd lmc mc mc m} t1        | j2                  |j4                        r| j7                          y y y rr   )r9   setupstatefnmodelr*   FITTING_layer_syncapplyr=   convert_modulemodel_to_deviceconfigure_ddpsetup_optimizersr"   setup_precision_pluginr   
optimizersrl   >torch.distributed.algorithms.ddp_comm_hooks.post_localSGD_hookdistributed
algorithmsddp_comm_hookspost_localSGD_hook
isinstancerV   PostLocalSGDState_enable_model_averaging)r_   trainer
trainer_fnpost_localSGDs       r`   r   zDDPStrategy.setup   s6   +++w']]%%
zz%%%***t/?/?))//

;DJ,,TZZ8***  !!'*  

+##%***!$//43C3CDbb$..0O0OP,,. Q +ra   r   c                 Z   | j                         }t        j                  d| d| j                          |;t        j
                  j                  t        j
                  j                               n	t               }|5  t        d||d| j                  cddd       S # 1 sw Y   yxY w)z^Wraps the model into a :class:`~torch.nn.parallel.distributed.DistributedDataParallel` module.z&setting up DDP model with device ids: z
, kwargs: N)module
device_ids )
determine_ddp_device_idsrN   rO   rU   torchcudastreamStreamr   r   )r_   r   r   ctxs       r`   _setup_modelzDDPStrategy._setup_model   s     224
		::,jQUQaQaPbcd8B8Nejj

 1 1 34T_Ta*c%JcRVRbRbc SSs   ?B!!B*c                    t         j                  | j                  j                   d       t	                | j                          | j                         | _        | j                  J d| j                  i}t        r*| j                  j                  dk7  r| j                  nd |d<   t        | j                  | j                  fi | y )Nz: setting up distributed...rC   cpu	device_id)rN   rO   rP   rQ   r   set_world_ranks_get_process_group_backendr[   r;   r\   r   rl   rf   r   )r_   rH   s     r`   r   zDDPStrategy.setup_distributed   s    		T^^,,--HIJ&*&E&E&G#''333"+T]]!;#6:6F6F6K6Ku6T$"2"2Z^F;d668S8S^W]^ra   c                 H    | j                   xs t        | j                        S rj   )r[   r   rl   rg   s    r`   r   z&DDPStrategy._get_process_group_backend   s     **m.[\`\l\l.mmra   c                 >   | j                   q| j                   j                  | j                  | j                  z  | j                  z          | j                   j                  | j                  | j                  z         | j                  xt        _	        t        _	        y rj   )r;   set_global_rank	node_rankrt   rk   set_world_sizero   rx   r   rw   utils_rank_zero_onlyrg   s    r`   r   zDDPStrategy.set_world_ranks   sy    ##/$$44T^^dFXFX5X[_[j[j5jk$$33DNNTEWEW4WX ;?:J:JJ27ra   c                 6   t         j                  | j                  j                   d       | j                  j
                  dk(  rTt        | j                  t              sJ t        | j                  | j                  | j                  | j                         y y )Nz: registering ddp hooksr   )r   r>   r?   r@   )rN   rO   rP   rQ   rl   rf   r   r   r   r!   rV   rW   rX   rg   s    r`   _register_ddp_hookszDDPStrategy._register_ddp_hooks   s{    		T^^,,--DEF   F*djj*ABBB#jj#33"11!%!7!7	 +ra   c                 f   t         j                  | j                  j                   d       | j                  t        d      ddlm}m}m	} | j                  D ]e  }t        |t              r|j                  }t        st        ||      nd}t        |||f      s|sDt        d|j                  j                   d       | j                  J t         j"                  j$                  j&                  j(                  j+                  | j                  | j                  j,                        | _        y )	Nz.: reinitializing optimizers with post localSGDz\Post-localSGD algorithm is used, but model averaging period is not provided to DDP strategy.r   )DistributedOptimizerPostLocalSGDOptimizerZeroRedundancyOptimizerFzKCurrently model averaging cannot work with a distributed optimizer of type .)periodwarmup_steps)rN   rO   rP   rQ   rY   
ValueErrortorch.distributed.optimr   r   r   r   r   r    
_optimizerr   rV   r   r   r   model_averaging	averagersPeriodicModelAveragerstart_localSGD_iterrZ   )r_   r   r   r   	optimizeris_distributed_optimizers         r`   r   z#DDPStrategy._enable_model_averaging   s   		T^^,,--[\]''/n  	qpI)%78%00	Zez)=Q'Rkp$)&=?T%UVZr a **334A7  ) ##///$00;;KKUUkk//d>R>R>f>f  l  
ra   r   closurepl.LightningModulec                     t        	|   |||fi |}| j                  |S |j                  D cg c]  }|d   D ]  }|j                  |  }}}| j                  j                  t        |             |S c c}}w )aI  Performs the actual optimizer step.

        Args:
            optimizer: the optimizer performing the step
            closure: closure calculating the loss value
            model: reference to the model, optionally defining optimizer step related hooks
            **kwargs: Any extra arguments to ``optimizer.step``

        params)rL   optimizer_steprZ   param_groupsgradaverage_parametersiter)
r_   r   r   r   rH   optimizer_outputr   paramr   rP   s
            r`   r   zDDPStrategy.optimizer_step  s    " !71)WeVvV'##%.%;%;s%;Ex\a\f\f\r%%%;s//V= ts   A:A:c                    t         j                  | j                  j                   d       t	        | j
                  t        j                        sJ | j                  | j
                        | _        | j                          y )Nz%: configuring DistributedDataParallel)
rN   rO   rP   rQ   r   r   plLightningModuler   r   rg   s    r`   r   zDDPStrategy.configure_ddp  s]    		T^^,,--RST$**b&8&8999&&tzz2
  "ra   c                 d    | j                   j                  dk(  ry | j                   j                  gS )Nr   )rl   rf   indexrg   s    r`   r   z$DDPStrategy.determine_ddp_device_ids"  s.      E)  &&''ra   argsc                     t               sy t        j                  j                         dk(  r/t        j                  j	                  | j                                y t        j                  j	                          y )Nnccl)r   )r   r   r   get_backendbarrierr   )r_   r   rH   s      r`   r   zDDPStrategy.barrier'  sT    *,((*f4%%1N1N1P%Q%%'ra   objsrcc                     t               s|S |g}t        j                  j                  ||t        j
                         |d   S )Nr   r   )r   r   r   broadcast_object_list_groupWORLD)r_   r   r   s      r`   	broadcastzDDPStrategy.broadcast1  s<    *,Je//S/M1vra   closure_lossc                     t        | j                  t              sy| j                  J | j                  j                  st        | j                  |       yy)z.Run before precision plugin executes backward.N)r   r   r   lightning_moduleautomatic_optimizationr#   )r_   r   s     r`   pre_backwardzDDPStrategy.pre_backward:  sJ     $**&=>$$000$$;; \: <ra   c                     | j                   }| j                  }t        |t              sy ||j                  ry t        |dd      sy | j                  ry |j                  }|j                          d| _        y )Nstatic_graphFT)	r   r   r   r   r   getattrr^   reducer_delay_all_reduce)r_   r   r   lmr   s        r`   post_backwardzDDPStrategy.post_backwardC  sp     

""%!89:22une4++ --!!#+/(ra   c                     t         j                  | j                  j                   d| j                   d       | j
                  J | j
                  j                  | j                         y )Nz: moving model to device [z]...)rN   rO   rP   rQ   rl   r   torg   s    r`   r   zDDPStrategy.model_to_deviceX  sU    		T^^,,--GHXHXGYY]^_zz%%%

d&&'ra   tensorr   	reduce_opc                 B    t        |t              rt        |||      S |S )a  Reduces a tensor from several distributed processes to one aggregated tensor.

        Args:
            tensor: the tensor to sync and reduce
            group: the process group to gather results from. Defaults to all processes (world)
            reduce_op: the reduction operation. Defaults to 'mean'/'avg'.
                Can also be a string 'sum' to calculate the sum during reduction.

        Return:
            reduced value, except when the input was not a tensor the output remains is unchanged

        )r   )r   r   r   )r_   r   r   r   s       r`   reducezDDPStrategy.reduce^  s"      ff%)&%9MMra   strategy_registryc           
          d}|D ]  \  }}|j                  || d| d|         d}|D ]#  \  }}}|j                  || d| d| d||       % y )	N))ddpr8   )	ddp_spawnrE   )r0   rF   )r3   rF   z"DDP strategy with `start_method` '')descriptionrD   )) ddp_find_unused_parameters_falseFr8   )ddp_find_unused_parameters_trueTr8   )&ddp_spawn_find_unused_parameters_falseFrE   )%ddp_spawn_find_unused_parameters_trueTrE   )r1   FrF   )r2   TrF   )r4   FrF   )r5   TrF   z.DDP strategy with `find_unused_parameters` as z and `start_method` ')r   find_unused_parametersrD   )register)clsr   entriesnamerD   fups         r`   register_strategieszDDPStrategy.register_strategiesr  s    
 #*D,&&@aP)	 '  #*	
 (/#D#|&&LSEQfgsfttuv'*) '  (/ra   	exceptionc                      t        |dd       y )Nz>.*Expected to have finished reduction in the prior iteration.*ay  It looks like your LightningModule has parameters that were not used in producing the loss returned by training_step. If this is intentional, you must enable the detection of unused parameters in DDP, either by setting the string value `strategy='ddp_find_unused_parameters_true'` or by setting the flag in the strategy with `strategy=DDPStrategy(find_unused_parameters=True)`.)patternnew_messager+   )r_   r  s     r`   on_exceptionzDDPStrategy.on_exception  s    Tt			
ra   c                    t         j                  | j                  j                   d       | j                  }t        | j                  t              ri| j                  j                  sL| j                  j                         j                  d      r#t        d| j                  j                   d       || _        ||j                  u|j                  j                  j                  t        j                   k(  rD| j"                  r8| j                  J | j"                  j%                  | j                        | _        t&        | Q          y )Nz: tearing down strategycan_set_static_graphzyYour model can run with static graph optimizations. For future training runs, we suggest you pass `Trainer(..., strategy=z%(static_graph=True))` to enable them.)rN   rO   rP   rQ   r   r   r   r   r   _get_ddp_logging_datagetr.   _trainerr   r   r*   r   r   revertrL   teardown)r_   	pl_modulerP   s     r`   r  zDDPStrategy.teardown  s	   		T^^,,--DEF))	djj"9:::**tzz/O/O/Q/U/UVl/m448NN4K4K3LLqs
 #DJ ! "".""((++y/@/@@  ::)))))00<DJra   )rI   N)r   z
pl.TrainerrI   Nrj   )r   )Nmean)?rQ   
__module____qualname____doc__r   r	   listr   devicer   r   r$   objectr   intstrr   r   r   rM   propertyboolrh   r   rl   ro   setterrt   dictry   rB   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   classmethodr   r   BaseExceptionr  r  __classcell__)rP   s   @r`   r7   r7   D   sS   U @D9=<@0404+/,0/304/3'9HO#1;<#1 #4#56#1 &&89	#1
  -#1 #9-#1 !(#1  )#1 #8,#1 !)#1  (}#1 )$#1 DE#1 #1 
#1J    6U\\ 6  6 3   $3 $4 $ $ Vs V V aDcN a  a +x} + + ] ] ! ! / /< d& d-D d d	_nC nK
0 
 @D	   "c'"  2F:;<	 
   
   4#((49*= (
 (S (C (D ( ( Z c *   ; ;D ; ; 0& 0T 0 0( ( (
 gm%-c]FNuU]_bUbOcFd	 &  4E  $     D 

m 

 

 

  ra   r7   c                   H    e Zd Zededdddfd       Zededdddfd       Zy)rR   wrapper_moduleoriginal_moduler   rI   Nc                 N    t        |t              r|j                  sd|_        y y y )NFr   r   r   require_backward_grad_syncr_   r  r   s      r`   on_after_inner_forwardz-_DDPForwardRedirection.on_after_inner_forward  s)     n&=>GmGm8=N5 Hn>ra   c                 N    t        |t              r|j                  sd|_        y y y )NTr"  r$  s      r`   on_after_outer_forwardz-_DDPForwardRedirection.on_after_outer_forward  s'    n&=>GmGm8<N5 Hn>ra   )rQ   r  r  r   r   r%  r'  r   ra   r`   rR   rR     sV    >V >Nb >gk > > =V =Nb =gk = =ra   rR   )Slogging
contextlibr   datetimer   typingr   r   r   r   r	   r
   r   torch.distributed"lightning_utilities.core.rank_zeror   r   r   torch.nnr   torch.nn.parallel.distributedr   torch.optim.optimizerr   typing_extensionsr   lightning.pytorchpytorchr   lightning.fabric.pluginsr   r   5lightning.fabric.plugins.collectives.torch_collectiver   lightning.fabric.strategiesr   &lightning.fabric.utilities.distributedr   r   r   r   r   r   "lightning.fabric.utilities.importsr   r   $lightning.fabric.utilities.optimizerr   lightning.fabric.utilities.seedr    lightning.fabric.utilities.typesr    lightning.pytorch.core.optimizerr    'lightning.pytorch.overrides.distributedr!   r"   r#   #lightning.pytorch.plugins.precisionr$   &lightning.pytorch.strategies.launchersr%   r&   %lightning.pytorch.strategies.parallelr'   %lightning.pytorch.strategies.strategyr(   r)    lightning.pytorch.trainer.statesr*   &lightning.pytorch.utilities.exceptionsr,   %lightning.pytorch.utilities.rank_zeror-   r.   6torch.distributed.algorithms.model_averaging.averagersr/   	getLoggerrQ   rN   _DDP_FORK_ALIASESr7   rR   r   ra   r`   <module>rH     s     "  I I   U   A + &  E T 9  C T F 6 5 ? v v 9 f B Q 6 C g gTg! x" xv=0 =ra   