
      i+                     p   d dl mZmZ d dlmZ d dl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 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,m-Z-m.Z.m/Z/m0Z0 d dl+m1Z2 d dl3m4Z4 d dl5mZ dZ6 G d de%      Z7 G d de*      Z8y)    )AbstractContextManagernullcontext)	timedelta)AnyLiteralOptionalUnionN)rank_zero_only)Tensor)Module)DistributedDataParallel)override)Accelerator)default_pg_timeout)ClusterEnvironment)CheckpointIO)	Precision)_MultiProcessingLauncher)_SubprocessScriptLauncher)ParallelStrategy)_StrategyRegistry)
TBroadcast_BackwardSyncControl)ReduceOp_distributed_is_initialized-_get_default_process_group_backend_for_device_init_dist_connection_sync_ddp_if_availablegroup)_TORCH_GREATER_EQUAL_2_3)ddp_forkddp_notebookc                       e Zd ZdZddddddedfdee   deeej                        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edej                  fd              Zedefd       Zej.                  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edefd       Z ededdfd       Z!e	 d0de"dee   dee#e$ef      de"fd       Z%ed ededdfd!       Z&ed1d"e'd#ede'fd$       Z(ededeee#ee"f   f   f fd%       Z)e	 d2ded&eee#ee"f   f   d'e*ddf fd(       Z+e,ed)e-ddfd*              Z.d/d+Z/defd,Z0d/d-Z1deee      fd.Z2 xZ3S )3DDPStrategyzKStrategy for multi-process single-device training on one or multiple nodes.Npopenacceleratorparallel_devicescluster_environmentcheckpoint_io	precisionprocess_group_backendtimeoutstart_method)r&   spawnfork
forkserverkwargsreturnc	                     t         
|   |||||       d| _        || _        || _        || _        t               | _        |	| _        y )N)r'   r(   r)   r*   r+      )	super__init__
_num_nodes_process_group_backend_timeout_start_method_DDPBackwardSyncControl_backward_sync_control_ddp_kwargs)selfr'   r(   r)   r*   r+   r,   r-   r.   r2   	__class__s             t/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/lightning/fabric/strategies/ddp.pyr7   zDDPStrategy.__init__8   sY     	#- 3' 	 	
 5J#-4)&=&?#!    c                 P    | j                   J | j                   | j                     S N)r(   
local_rankr?   s    rA   root_devicezDDPStrategy.root_deviceR   s+     $$000$$T__55rB   c                     | j                   S rD   r8   rF   s    rA   	num_nodeszDDPStrategy.num_nodesX   s    rB   rJ   c                     || _         y rD   rI   )r?   rJ   s     rA   rJ   zDDPStrategy.num_nodes\   s     $rB   c                 H    | j                   t        | j                         S dS )Nr   )r(   lenrF   s    rA   num_processeszDDPStrategy.num_processesa   s$    -1-B-B-Ns4(()UTUUrB   c                 N    | j                   | j                  z  | j                  dS )N)num_replicasrank)rJ   rN   global_rankrF   s    rA   distributed_sampler_kwargsz&DDPStrategy.distributed_sampler_kwargse   s$     "&$2D2D!DtO_O_``rB   c                     | j                   S rD   )r9   rF   s    rA   r,   z!DDPStrategy.process_group_backendj   s    ***rB   c                     | j                   J | j                  dk(  r1t        | j                   | j                  | j                        | _        y t        | | j                        | _        y )Nr&   )r.   )r)   r;   r   rN   rJ   	_launcherr   rF   s    rA   _configure_launcherzDDPStrategy._configure_launchern   sZ    ''333(6t7O7OQUQcQceiesestDN5dI[I[\DNrB   c                 B    t         |           | j                          y rD   )r6   setup_environment_setup_distributed)r?   r@   s    rA   rY   zDDPStrategy.setup_environmentv   s    !#!rB   modulec                    | 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.N)r[   
device_ids )_determine_ddp_device_idstorchcudastreamStreamr   r   r>   )r?   r[   r]   ctxs       rA   setup_modulezDDPStrategy.setup_module{   sd     335
8B8Nejj

 1 1 34T_Ta*d&ZdSWScScd SSs   A<<Bc                 :    |j                  | j                         y rD   )torG   )r?   r[   s     rA   module_to_devicezDDPStrategy.module_to_device   s    		$""#rB   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

        )rj   )
isinstancer   r   )r?   ri   r    rj   s       rA   
all_reducezDDPStrategy.all_reduce   s"      ff%)&%9MMrB   argsc                    t               sy t        j                  j                         dk(  r/t        j                  j	                  | j                                y 	 t        j                  j	                          y # t        $ rY}dt        |      v rAt        j                  d| j                        }t        j                  j                  |       n Y d }~y d }~ww xY w)Nnccl)r]   PrivateUse1HooksInterfaceg        )device)r   r`   distributedget_backendbarrierr_   RuntimeErrorstrri   rG   rm   )r?   rn   r2   edummy_tensors        rA   ru   zDDPStrategy.barrier   s    *,((*f4%%1O1O1Q%R	!!))+ .#a&8 $)<<D<L<L#ML%%00> ?s   A< <	CACCobjsrcc                     t               s|S |g}t        j                  j                  ||t        j
                         |d   S )Nr   r   )r   r`   rs   broadcast_object_list_groupWORLD)r?   rz   r{   s      rA   	broadcastzDDPStrategy.broadcast   s<    *,Je//S/M1vrB   c                 Z    t        |t              r|j                  }t        |   |      S rD   )rl   r   r[   r6   get_module_state_dict)r?   r[   r@   s     rA   r   z!DDPStrategy.get_module_state_dict   s'    f56]]Fw,V44rB   
state_dictstrictc                 b    t        |t              r|j                  }t        |   |||       y )N)r[   r   r   )rl   r   r[   r6   load_module_state_dict)r?   r[   r   r   r@   s       rA   r   z"DDPStrategy.load_module_state_dict   s.     f56]]F&fTZ&[rB   strategy_registryc                 z    d}|D ]  \  }}|j                  || d|d|         |j                  d| ddd	       y )
N))ddpr&   )	ddp_spawnr/   )r"   r0   )r#   r0   z DDP strategy with `start_method=`)descriptionr.   ddp_find_unused_parameters_truezBAlias for `find_unused_parameters_true` and `start_method='popen'`Tr&   )r   find_unused_parametersr.   )register)clsr   entriesnamer.   s        rA   register_strategieszDDPStrategy.register_strategies   sg    
 #*D,&&>|>NaP)	 '  #* 	""-\#'  	# 	
rB   c                 (   | j                          | j                         | _        | j                  J d| j                  i}t
        r*| j                  j                  dk7  r| j                  nd |d<   t        | j                  | j                  fi | y )Nr-   cpu	device_id)	_set_world_ranks_get_process_group_backendr9   r)   r:   r!   rG   typer   )r?   r2   s     rA   rZ   zDDPStrategy._setup_distributed   s    &*&E&E&G#''333"+T]]!;#6:6F6F6K6Ku6T$"2"2Z^F;d668S8S^W]^rB   c                 H    | j                   xs t        | j                        S rD   )r9   r   rG   rF   s    rA   r   z&DDPStrategy._get_process_group_backend   s     **m.[\`\l\l.mmrB   c                 >   | j                   q| j                   j                  | j                  | j                  z  | j                  z          | j                   j                  | j                  | j                  z         | j                  xt        _	        t        _	        y rD   )r)   set_global_rank	node_rankrN   rE   set_world_sizerJ   rR   r
   rQ   utils_rank_zero_onlyrF   s    rA   r   zDDPStrategy._set_world_ranks   sy    ##/$$44T^^dFXFX5X[_[j[j5jk$$33DNNTEWEW4WX ;?:J:JJ27rB   c                 f    | j                   j                  dk(  rd S | j                   j                  gS )Nr   )rG   r   indexrF   s    rA   r_   z%DDPStrategy._determine_ddp_device_ids   s/    '',,5tSD<L<L<R<R;SSrB   )r3   N)Nmean)r   )T)4__name__
__module____qualname____doc__r   r   r   listr`   rr   r   r   r   rw   r   r   r   r7   propertyr   rG   intrJ   setterrN   dictrS   r,   rW   rY   r   r   re   rh   r   r	   r   rm   ru   r   r   r   boolr   classmethodr   r   rZ   r   r   r_   __classcell__)r@   s   @rA   r%   r%   5   s   U .29=<@04)-/3'9HO"k*" #4#56" &&89	"
  -" I&"  (}" )$" DE" " 
"4 6U\\ 6  6 3   $3 $4 $ $ Vs V V aDcN a  a +x} + + ] ] " " e6 e.E e e $v $$ $ $ gm%-c]FNuU]_bUbOcFd	 & S C D  $ Z c *   5F 5tCsF{AS<S7T 5 5
 X\\\*.sE#v+4F/F*G\QU\	\ \ 
4E 
$ 
  
,_nC nKT8DI+> TrB   r%   c                   *    e Zd Zedededefd       Zy)r<   r[   enabledr3   c                     |s
t               S t        |t              s:t        d| j                  j
                   d|j                  j
                   d      |j                         S )z{Blocks gradient synchronization inside the :class:`~torch.nn.parallel.distributed.DistributedDataParallel`
        wrapper.zABlocking backward sync is only possible if the module passed to `zA.no_backward_sync` is wrapped in `DistributedDataParallel`. Got: .)r   rl   r   	TypeErrorr@   r   no_sync)r?   r[   r   s      rA   no_backward_syncz(_DDPBackwardSyncControl.no_backward_sync   si     = &"9:^^,,- .))22316 
 ~~rB   N)r   r   r   r   r   r   r   r   r^   rB   rA   r<   r<      s*     v    AW    rB   r<   )9
contextlibr   r   datetimer   typingr   r   r   r	   r`   torch.distributed"lightning_utilities.core.rank_zeror
   r   r   torch.nnr   torch.nn.parallel.distributedr   typing_extensionsr   )lightning.fabric.accelerators.acceleratorr   5lightning.fabric.plugins.collectives.torch_collectiver   9lightning.fabric.plugins.environments.cluster_environmentr   )lightning.fabric.plugins.io.checkpoint_ior   "lightning.fabric.plugins.precisionr   5lightning.fabric.strategies.launchers.multiprocessingr   7lightning.fabric.strategies.launchers.subprocess_scriptr   $lightning.fabric.strategies.parallelr   $lightning.fabric.strategies.registryr   $lightning.fabric.strategies.strategyr   r   &lightning.fabric.utilities.distributedr   r   r   r   r   r    r~   "lightning.fabric.utilities.importsr!   $lightning.fabric.utilities.rank_zero_DDP_FORK_ALIASESr%   r<   r^   rB   rA   <module>r      s    ;  0 0   U   A & A T X B 8 Z ] A B Q  C G ? ~T" ~TB 2  rB   