a
    d                     @   s   d dl mZ d dlmZmZmZm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 G dd deeZdS )    )ABC)AnyDictListOptionalN)Tensor)Accelerator)ClusterEnvironment)CheckpointIO)	Precision)Strategy_all_gather_ddp_if_available)ReduceOpc                       s>  e Zd ZdZd!ee eeej  ee	 ee
 ee d f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ef  dddZd"eee eedddZd#eeedddZdd fdd Z  ZS )$ParallelStrategyz:Strategy for training with multiple processes in parallel.N)acceleratorparallel_devicescluster_environmentcheckpoint_io	precisionc                    s"   t  j|||d || _|| _d S )N)r   r   r   )super__init__r   r   )selfr   r   r   r   r   	__class__ m/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/lightning_fabric/strategies/parallel.pyr       s    zParallelStrategy.__init__)returnc                 C   s   | j d ur| j  S dS Nr   )r   global_rankr   r   r   r   r   ,   s    zParallelStrategy.global_rankc                 C   s   | j d ur| j  S dS r   )r   
local_rankr    r   r   r   r!   0   s    zParallelStrategy.local_rankc                 C   s   | j d ur| j  S dS r   )r   	node_rankr    r   r   r   r"   4   s    zParallelStrategy.node_rankc                 C   s   | j d ur| j  S dS )N   )r   
world_sizer    r   r   r   r$   8   s    zParallelStrategy.world_sizec                 C   s
   | j dkS r   )r   r    r   r   r   is_global_zero<   s    zParallelStrategy.is_global_zeroc                 C   s   | j S NZ_parallel_devicesr    r   r   r   r   @   s    z!ParallelStrategy.parallel_devices)r   r   c                 C   s
   || _ d S r&   r'   )r   r   r   r   r   r   D   s    c                 C   s   | j | jdS )zArguments for the ``DistributedSampler``.

        If this method is not defined, or it returns ``None``, then the ``DistributedSampler`` will not be used.
        )Znum_replicasZrank)r$   r   r    r   r   r   distributed_sampler_kwargsH   s    z+ParallelStrategy.distributed_sampler_kwargsF)tensorgroup
sync_gradsr   c                 C   s   t |||dS )z&Perform a all_gather on all processes.)r*   r+   r   )r   r)   r*   r+   r   r   r   
all_gatherP   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)	torchr)   intZroot_deviceZ
all_reducer   ZSUMboolr$   )r   r-   r.   r   r   r   reduce_boolean_decisionT   s    z(ParallelStrategy.reduce_boolean_decisionc                    s"   | j d usJ | j   t  S r&   )r   teardownr   r    r   r   r   r4   e   s    
zParallelStrategy.teardown)NNNNN)NF)T) __name__
__module____qualname____doc__r   r   r   r0   r/   r	   r
   r   r   propertyr1   r   r!   r"   r$   r2   r%   r   setterr   strr   r(   r   r,   r3   r4   __classcell__r   r   r   r   r      s@        r   )abcr   typingr   r   r   r   r0   r   Z)lightning_fabric.accelerators.acceleratorr   Z9lightning_fabric.plugins.environments.cluster_environmentr	   Z)lightning_fabric.plugins.io.checkpoint_ior
   Z"lightning_fabric.plugins.precisionr   Z$lightning_fabric.strategies.strategyr   Z&lightning_fabric.utilities.distributedr   Z lightning_fabric.utilities.typesr   r   r   r   r   r   <module>   s   