a
    d                     @   s   d dl mZmZ d dlmZ d dlmZmZmZm	Z	m
Z
 d dlZd dlmZ d dlZd dlmZmZ d dlmZmZ d dlmZ d d	lmZ d d
lmZ G dd deeZdS )    )ABCabstractmethod)contextmanager)AnyDict	GeneratorListOptionalN)Tensor)CheckpointIOClusterEnvironment)_all_gather_ddp_if_availableReduceOp)	LayerSync)PrecisionPlugin)Strategyc                       st  e Zd ZdZd)ed eeej  ee ee	 ee
 d fddZeeejddd	Zeedd
dZeedddZeedddZeedddZeedddZeeeej  dddZejeeej  ddddZeeeef dddZeddddZd*eee eeddd Zd+eeed"d#d$Ze e!dd%d&Z"dd fd'd(Z#  Z$S ),ParallelStrategyz8Plugin for training with multiple processes in parallel.Nzpl.accelerators.Accelerator)acceleratorparallel_devicescluster_environmentcheckpoint_ioprecision_pluginc                    s(   t  j|||d || _|| _d | _d S )N)r   r   r   )super__init__r   r   Z_layer_sync)selfr   r   r   r   r   	__class__ n/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/strategies/parallel.pyr       s    zParallelStrategy.__init__)returnc                 C   s   dS )zReturn the root device.Nr   r   r   r   r   root_device-   s    zParallelStrategy.root_devicec                 C   s   | j d ur| j  S dS Nr   )r   global_rankr    r   r   r   r#   2   s    zParallelStrategy.global_rankc                 C   s   | j d ur| j  S dS r"   )r   
local_rankr    r   r   r   r$   6   s    zParallelStrategy.local_rankc                 C   s   | j d ur| j  S dS r"   )r   	node_rankr    r   r   r   r%   :   s    zParallelStrategy.node_rankc                 C   s   | j d ur| j  S dS )N   )r   
world_sizer    r   r   r   r'   >   s    zParallelStrategy.world_sizec                 C   s
   | j dkS r"   )r#   r    r   r   r   is_global_zeroB   s    zParallelStrategy.is_global_zeroc                 C   s   | j S NZ_parallel_devicesr    r   r   r   r   F   s    z!ParallelStrategy.parallel_devices)r   r   c                 C   s
   || _ d S r)   r*   )r   r   r   r   r   r   J   s    c                 C   s"   t | jd urt| jnd| jdS )Nr   )Znum_replicasZrank)dictr   lenr#   r    r   r   r   distributed_sampler_kwargsN   s    z+ParallelStrategy.distributed_sampler_kwargs)tracer   c                 C   s   dS )z/Function to re-conciliate processes on failure.Nr   )r   r.   r   r   r   reconciliate_processesU   s    z'ParallelStrategy.reconciliate_processesF)tensorgroup
sync_gradsr   c                 C   s   t |||dS )z&Perform a all_gather on all processes.)r1   r2   )r   )r   r0   r1   r2   r   r   r   
all_gatherX   s    zParallelStrategy.all_gatherT)decisionallr   c                 C   sB   t jt|| jd}| j|tjd}|r6t|| jknt|}|S )a  Reduces a boolean decision over distributed processes. By default is analagous to ``all`` from the
        standard library, returning ``True`` only if all input decisions evaluate to ``True``. If ``all`` is set to
        ``False``, it behaves like ``any`` instead.

        Args:
            decision: A single input decision.
            all: Whether to logically emulate ``all`` or ``any``. Defaults to True.

        Returns:
            bool: The reduced boolean decision.
        )device)Z	reduce_op)	torchr0   intr!   reducer   ZSUMboolr'   )r   r4   r5   r   r   r   reduce_boolean_decision\   s    z(ParallelStrategy.reduce_boolean_decisionc                 c   sN   t | jtjjjrD| j  dV  W d   qJ1 s80    Y  ndV  dS )zBlocks ddp sync gradients behaviour on backwards pass.

        This is useful for skipping sync when accumulating gradients, reducing communication overhead
        Returns: context manager with sync behaviour off
        N)
isinstancemodelplZ	utilitiestypesZDistributedDataParallelZno_syncr    r   r   r   block_backward_syncm   s    &z$ParallelStrategy.block_backward_syncc                    s&   | j d usJ | j   t   d S r)   )r   teardownr   r    r   r   r   rA   z   s    
zParallelStrategy.teardown)NNNNN)NF)T)%__name__
__module____qualname____doc__r	   r   r7   r6   r   r   r   r   propertyr   r!   r8   r#   r$   r%   r'   r:   r(   r   setterr   strr   r-   r/   r
   r3   r;   r   r   r@   rA   __classcell__r   r   r   r   r      sL        r   )abcr   r   
contextlibr   typingr   r   r   r   r	   r7   r
   Zpytorch_lightningr>   Zlightning_fabric.pluginsr   r   Z&lightning_fabric.utilities.distributedr   r   Zpytorch_lightning.pluginsr   Z#pytorch_lightning.plugins.precisionr   Z%pytorch_lightning.strategies.strategyr   r   r   r   r   r   <module>   s   