a
    d"                     @   sL  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 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+ d dl'm,Z- d dl'm.Z. d dl/m0Z0 dZ1G dd de#Z2G dd de%Z3dS )    )contextmanager)	timedelta)AnyDict	GeneratorListOptionalUnionN)Tensor)Module)DistributedDataParallel)Literal)Accelerator)default_pg_timeout)ClusterEnvironment)CheckpointIO)	Precision)_MultiProcessingLauncher)_SubprocessScriptLauncher)ParallelStrategy)_BackwardSyncControl
TBroadcast)_distributed_available-_get_default_process_group_backend_for_device_init_dist_connection_sync_ddp_if_availablegroup)ReduceOp)rank_zero_only)ddp_forkddp_notebookc                       s  e Zd ZdZddddddedfee eeej	  ee
 ee ee ee ee ed edd
 fddZeej	dd	d
ZeedddZejeddddZeedddZeeeef dddZeee dddZddddZdd fddZeedddZeddddZ d4e!ee ee"e#ef  e!dd d!Z$eedd"d#d$Z%d5e&ee&d&d'd(Z'e(edd)d*d+Z)ddd,d-Z*edd.d/Z+ddd0d1Z,eee  dd2d3Z-  Z.S )6DDPStrategyzKStrategy for multi-process single-device training on one or multiple nodes.Npopen)r#   spawnforkZ
forkserver)
acceleratorparallel_devicescluster_environmentcheckpoint_io	precisionprocess_group_backendtimeoutstart_methodkwargsreturnc	           
         s@   t  j|||||d d| _|| _|| _|| _t | _|	| _d S )N)r&   r'   r(   r)   r*      )	super__init__
_num_nodes_process_group_backend_timeout_start_method_DDPBackwardSyncControlZ_backward_sync_control_ddp_kwargs)
selfr&   r'   r(   r)   r*   r+   r,   r-   r.   	__class__ h/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/lightning_fabric/strategies/ddp.pyr2   5   s    zDDPStrategy.__init__)r/   c                 C   s   | j d usJ | j | j S N)r'   
local_rankr9   r<   r<   r=   root_deviceO   s    zDDPStrategy.root_devicec                 C   s   | j S r>   r3   r@   r<   r<   r=   	num_nodesT   s    zDDPStrategy.num_nodes)rC   r/   c                 C   s
   || _ d S r>   rB   )r9   rC   r<   r<   r=   rC   X   s    c                 C   s   | j d urt| j S dS )Nr   )r'   lenr@   r<   r<   r=   num_processes]   s    zDDPStrategy.num_processesc                 C   s   t | j| j | jdS )N)Znum_replicasrank)dictrC   rE   global_rankr@   r<   r<   r=   distributed_sampler_kwargsa   s    z&DDPStrategy.distributed_sampler_kwargsc                 C   s   | j S r>   )r4   r@   r<   r<   r=   r+   e   s    z!DDPStrategy.process_group_backendc                 C   sB   | j d usJ | jdkr.t| j | j| j| _nt| | jd| _d S )Nr#   )r-   )r(   r6   r   rE   rC   Z	_launcherr   r@   r<   r<   r=   _configure_launcheri   s    
zDDPStrategy._configure_launcherc                    s   |    t   d S r>   )_setup_distributedr1   setup_environmentr@   r:   r<   r=   rL   p   s    zDDPStrategy.setup_environmentmoduler/   c                 C   s   t f ||  d| jS )z^Wraps the model into a :class:`~torch.nn.parallel.distributed.DistributedDataParallel` module.)rN   
device_ids)r   _determine_ddp_device_idsr8   r9   rN   r<   r<   r=   setup_modulet   s    zDDPStrategy.setup_modulec                 C   s   | | j d S r>   )torA   rQ   r<   r<   r=   module_to_devicex   s    zDDPStrategy.module_to_devicemean)tensorr   	reduce_opr/   c                 C   s   t |trt|||d}|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
        )rW   )
isinstancer
   r   )r9   rV   r   rW   r<   r<   r=   
all_reduce{   s    
zDDPStrategy.all_reduce)argsr.   r/   c                 O   s:   t  s
d S tj dkr,tjj|  d n
tj  d S )NZnccl)rO   )r   torchdistributedget_backendbarrierrP   )r9   rZ   r.   r<   r<   r=   r^      s
    zDDPStrategy.barrierr   )objsrcr/   c                 C   s<   t  s
|S |g}| j|kr d g}tjj||tjd |d S )Nr   r   )r   rH   r[   r\   Zbroadcast_object_list_groupZWORLD)r9   r_   r`   r<   r<   r=   	broadcast   s    
zDDPStrategy.broadcast)strategy_registryr/   c                 C   s0   d}|D ]"\}}|j || d|d|d qd S )N))Zddpr#   )Z	ddp_spawnr$   )r    r%   )r!   r%   z DDP strategy with `start_method=`)descriptionr-   )register)clsrc   entriesnamer-   r<   r<   r=   register_strategies   s    
zDDPStrategy.register_strategiesc                 C   s@   |    | jt_|  | _| jd us(J t| j| j| jd d S )N)r,   )	_set_world_ranksrH   r   rF   _get_process_group_backendr4   r(   r   r5   r@   r<   r<   r=   rK      s
    
zDDPStrategy._setup_distributedc                 C   s   | j pt| jS r>   )r4   r   rA   r@   r<   r<   r=   rl      s    z&DDPStrategy._get_process_group_backendc                 C   sL   | j d u rd S | j | j| j | j  | j | j| j  | j  t_	d S r>   )
r(   Zset_global_rankZ	node_rankrE   r?   Zset_world_sizerC   rH   r   rF   r@   r<   r<   r=   rk      s
    
zDDPStrategy._set_world_ranksc                 C   s   | j jdkrd S | j jgS )Ncpu)rA   typeindexr@   r<   r<   r=   rP      s    z%DDPStrategy._determine_ddp_device_ids)NrU   )r   )/__name__
__module____qualname____doc__r   r   r   r   r[   Zdevicer   r   r   strr   r   r   r2   propertyrA   intrC   setterrE   r   rI   r+   rJ   rL   r   r   rR   rT   r
   r	   r   rY   r^   r   rb   classmethodrj   rK   rl   rk   rP   __classcell__r<   r<   r:   r=   r"   2   sd    	r"   c                   @   s    e Zd ZeeedddZdS )r7   rM   c                 c   sZ   t |ts(td| jj d|jj d|  dV  W d   n1 sL0    Y  d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: .N)rX   r   	TypeErrorr;   rp   Zno_syncrQ   r<   r<   r=   no_backward_sync   s    

z(_DDPBackwardSyncControl.no_backward_syncN)rp   rq   rr   r   r   r   r|   r<   r<   r<   r=   r7      s   r7   )4
contextlibr   datetimer   typingr   r   r   r   r   r	   r[   Ztorch.distributedr
   Ztorch.nnr   Ztorch.nn.parallel.distributedr   Ztyping_extensionsr   Z)lightning_fabric.accelerators.acceleratorr   Z5lightning_fabric.plugins.collectives.torch_collectiver   Z9lightning_fabric.plugins.environments.cluster_environmentr   Z)lightning_fabric.plugins.io.checkpoint_ior   Z"lightning_fabric.plugins.precisionr   Z5lightning_fabric.strategies.launchers.multiprocessingr   Z7lightning_fabric.strategies.launchers.subprocess_scriptr   Z$lightning_fabric.strategies.parallelr   Z$lightning_fabric.strategies.strategyr   r   Z&lightning_fabric.utilities.distributedr   r   r   r   r   ra   r   Z$lightning_fabric.utilities.rank_zeror   Z_DDP_FORK_ALIASESr"   r7   r<   r<   r<   r=   <module>   s2     