a
    dD3                     @   s  d dl Z d dlZd dlZd dlmZ d dlmZmZmZm	Z	m
Z
mZmZmZ d dlZd dlm  mZ d dlmZ d dlmZ d dlmZmZmZ d dlmZ d dlmZ d d	lmZ d d
l m!Z! ej"# rd dl$m%Z% nG dd dZ%e &e'Z(d0ee
e e	e dddZ)eee*e	e dddZ+e,dddZ-d1ee
e e
ee!e.f  edddZ/d2ee
e e
ee!e.f  edddZ0G dd dej1j2Z3eeeddd Z4d3ee
d" e,ed#d$d%Z5d4ee.e
e* e
e* edd&d'd(Z6ej7e.d)d*d+Z8G d,d- d-eZ9G d.d/ d/eZ:dS )5    N)nullcontext)AnyIterableIteratorListOptionalSizedTupleUnion)module_available)Tensor)DatasetDistributedSamplerSampler)ClusterEnvironment)_TORCH_GREATER_EQUAL_1_12)rank_zero_info)ReduceOpgroupc                   @   s   e Zd ZdZdS )r   N)__name__
__module____qualname__WORLD r   r   o/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/lightning_fabric/utilities/distributed.pyr      s   r   )resultr   returnc                    s`  |du rt jjj}|  } t j|}t jj|d | jdkrJt| ||S t j	| j
| jd  fddt|D }t jj| |d t |jddjtfdd	|D }|rt| ||S g }    }t|D ]}|d ||  qt| |fd
dt|D }t j|| t|D ](\}	}
dd |
D }||	 | ||	< q2|S )ah  Function to gather all tensors from several DDP processes onto a list that is broadcasted to all processes.

    Works on tensors that have the same number of dimensions, but where each dimension may differ. In this case
    tensors are padded, gathered and then trimmed to secure equal workload for all processes.

    Args:
        result: The value to sync
        group: The process group to gather results from. Defaults to all processes (world)

    Return:
        gathered_result: List with size equal to the process group where
            gathered_result[i] corresponds to result tensor from process i
    Nr   r   devicec                    s   g | ]}t  qS r   torchZ
zeros_like.0_)
local_sizer   r   
<listcomp>:       z'_gather_all_tensors.<locals>.<listcomp>Zdimc                 3   s   | ]}t | kV  qd S N)all)r#   Zls)max_sizer   r   	<genexpr>=   r'   z&_gather_all_tensors.<locals>.<genexpr>c                    s   g | ]}t  qS r   r    r"   )result_paddedr   r   r&   J   r'   c                 S   s   g | ]}t |qS r   )slice)r#   Zdim_sizer   r   r   r&   M   r'   )r!   distributedr   r   
contiguousget_world_sizebarrierndim_simple_gather_all_tensorstensorshaper   range
all_gatherstackmaxvaluesr*   detachcpureversedappenditemFpad	enumerate)r   r   
world_sizeZlocal_sizesZall_sizes_equalZpad_dimsZpad_byvalgathered_resultidxZ	item_sizeZslice_paramr   )r%   r+   r-   r   _gather_all_tensors   s4    


rH   )r   r   rD   r   c                    s*    fddt |D }tj| | |S )Nc                    s   g | ]}t  qS r   r    r"   r   r   r   r&   S   r'   z._simple_gather_all_tensors.<locals>.<listcomp>)r7   r!   r/   r8   )r   r   rD   rF   r   rI   r   r4   R   s    r4   r   c                  C   s&   ddl m}  tj r tj p$|  S )Nr   _tpu_distributed)Z!lightning_fabric.accelerators.tpurL   r!   r/   is_availableis_initializedrK   r   r   r   _distributed_availableX   s    rO   )r   r   	reduce_opr   c                 C   s   t  rt| ||dS | S )a  Function to reduce a tensor across worker processes during distributed training.

    Args:
        result: The value to sync and reduce (typically tensor or number)
        group: The process group to gather results from. Defaults to all processes (world)
        reduce_op: The reduction operation. Defaults to sum.
            Can also be a string of 'avg', 'mean' to calculate the mean during reduction.

    Return:
        reduced value
    )r   rP   )rO   	_sync_ddp)r   r   rP   r   r   r   _sync_ddp_if_available^   s    rR   c                 C   s   d}|du rt jjj}t|trH| dv r8tj}d}qLt	t|
 }n|}tdrddlm} | rtjdd	kr|  d
v rtd |  } t jj|d t jj| ||dd |r| t j| } | S )a  Function to reduce the tensors from several DDP processes to one main process.

    Args:
        result: The value to sync and reduce (typically tensor or number)
        group: The process group to gather results from. Defaults to all processes (world)
        reduce_op: The reduction operation. Defaults to sum.
            Can also be a string of 'avg', 'mean' to calculate the mean during reduction.

    Return:
        reduced value
    FN)avgmeanTz,habana_frameworks.torch.utils.library_loaderr   )is_habana_avaialbleZHCCL_DISTRIBUTED_BACKEND1)ztorch.LongTensorztorch.hpu.LongTensorz0Long tensor unsupported on HPU, casting to floatr   )opr   async_op)r!   r/   r   r   
isinstancestrlowerr   SUMgetattrupperr   Z,habana_frameworks.torch.utils.library_loaderrU   osenvirongettyper   floatr2   
all_reducer1   )r   r   rP   Zdivide_by_world_sizerW   rU   r   r   r   rQ   q   s0    


rQ   c                   @   sL   e Zd Zeejfeeed edddZ	eeee
edf dddZdS )	
_AllGathertorch.distributed.ProcessGroup)ctxr5   r   r   c                    sJ   || _  fddttjj|dD }tjj| |d tj|dd}|S )Nc                    s   g | ]}t  qS r   r    r"   r5   r   r   r&      r'   z&_AllGather.forward.<locals>.<listcomp>r   r   r(   )r   r7   r!   r/   r1   r8   r9   )rg   r5   r   Zgathered_tensorr   rh   r   forward   s
     z_AllGather.forwardN)rg   grad_outputr   c                 G   s8   t |}t jj|t jjjd| jd |t j  d fS )NF)rW   rX   r   )r!   catr/   rd   r   r\   r   Zget_rank)rg   rj   r   r   r   backward   s    
z_AllGather.backward)r   r   r   staticmethodr   r   r   r   r   ri   r	   rl   r   r   r   r   re      s   re   )r5   r   r   c                 C   s4   t jdkrtst| |S ddl}|jjj	| |S )z!Compatibility layer with Windows.win32r   N)
sysplatformr   re   applyZtorch.distributed.nnr/   nn
functionalr8   )r5   r   r!   r   r   r   _functional_all_gather   s    rt   Frf   )r5   r   
sync_gradsr   c                 C   sX   t  s
| S |  } |rt nt  t| |}W d   n1 sD0    Y  t|S )am  Function to gather a tensor from several distributed processes.

    Args:
        tensor: Tensor of shape (batch, ...)
        group: The process group to gather results from. Defaults to all processes (world)
        sync_grads: Flag that allows users to synchronize gradients for all_gather op

    Return:
        A tensor of shape (world_size, batch, ...)
    N)rO   r0   r   r!   Zno_gradrt   r9   )r5   r   ru   Zgathered_tensorsr   r   r   _all_gather_ddp_if_available   s    (rv   )cluster_environmenttorch_distributed_backendglobal_rankrD   kwargsr   c              	   K   s   t j stdt j r*td dS |dur6|n|  }|durJ|n|  }| j	t
jd< t| jt
jd< td| d|d  d	|  t jj|f||d
| td d| d| dd d dS )a  Utility function to initialize distributed connection by setting env variables and initializing the
    distributed process group.

    Args:
        cluster_environment: ``ClusterEnvironment`` instance
        torch_distributed_backend: Backend to use (includes `nccl` and `gloo`)
        global_rank: Rank of the current process
        world_size: Number of processes in the group
        kwargs: Kwargs for ``init_process_group``

    Raises:
        RuntimeError:
            If ``torch.distributed`` is not available
    zOtorch.distributed is not available. Cannot initialize distributed process groupz7torch.distributed is already initialized. Exiting earlyNZMASTER_ADDRZMASTER_PORTz'Initializing distributed: GLOBAL_RANK: z
, MEMBER:    /)ZrankrD   zd----------------------------------------------------------------------------------------------------z
distributed_backend=z5
All distributed processes registered. Starting with z processes

)r!   r/   rM   RuntimeErrorrN   logdebugry   rD   Zmain_addressr_   r`   rZ   Z	main_portinfoZinit_process_groupr   )rw   rx   ry   rD   rz   r   r   r   _init_dist_connection   s(    


 r   )r   r   c                 C   s   | j dkrdS dS )NZcudaZncclZgloo)rb   r   r   r   r   -_get_default_process_group_backend_for_device   s    r   c                   @   sT   e Zd ZdZeeef ddddZee	dddZ
ed	d
dZdd	ddZdS )_DatasetSamplerWrapperz6Dataset to create indexes from `Sampler` or `Iterable`N)samplerr   c                 C   s:   t |tstdt|tdkr*td|| _d | _d S )Na  You seem to have configured a sampler in your DataLoader which does not provide `__len__` method. The sampler was about to be replaced by `DistributedSamplerWrapper` since `replace_sampler_ddp` is True and you are using distributed training. Either provide `__len__` method in your sampler, remove it from DataLoader or set `replace_sampler_ddp=False` if you want to handle distributed sampling yourself.infa  You seem to have configured a sampler in your DataLoader which does not provide finite `__len__` method. The sampler was about to be replaced by `DistributedSamplerWrapper` since `replace_sampler_ddp` is True and you are using distributed training. Either provide `__len__` method in your sampler which returns a finite number, remove it from DataLoader or set `replace_sampler_ddp=False` if you want to handle distributed sampling yourself.)rY   r   	TypeErrorlenrc   _sampler_sampler_list)selfr   r   r   r   __init__  s    
z_DatasetSamplerWrapper.__init__)indexr   c                 C   s    | j d u rt| j| _ | j | S r)   )r   listr   )r   r   r   r   r   __getitem__  s    
z"_DatasetSamplerWrapper.__getitem__rJ   c                 C   s
   t | jS r)   )r   r   r   r   r   r   __len__$  s    z_DatasetSamplerWrapper.__len__c                 C   s   t | j| _dS )z4Reset the sampler list in order to get new sampling.N)r   r   r   r   r   r   r   reset'  s    z_DatasetSamplerWrapper.reset)r   r   r   __doc__r
   r   r   r   intr   r   r   r   r   r   r   r   r     s
   r   c                       sF   e Zd ZdZeeef eedd fddZe	d fddZ
  ZS )	DistributedSamplerWrappera  Wrapper over ``Sampler`` for distributed training.

    Allows you to use any sampler in distributed mode. It will be automatically used by Lightning in distributed mode if
    sampler replacement is enabled.

    Note:
        The purpose of this wrapper is to take care of sharding the sampler indices. It is up to the underlying
        sampler to handle randomness and shuffling. The ``shuffle`` and ``seed`` arguments on this wrapper won't
        have any effect.
    N)r   argsrz   r   c                    s"   t  jt|g|R i | d S r)   )superr   r   )r   r   r   rz   	__class__r   r   r   8  s    z"DistributedSamplerWrapper.__init__rJ   c                    s"    j    fddt  D S )Nc                 3   s   | ]} j | V  qd S r)   )dataset)r#   r   r   r   r   r,   =  r'   z5DistributedSamplerWrapper.__iter__.<locals>.<genexpr>)r   r   r   __iter__r   r   r   r   r   ;  s    
z"DistributedSamplerWrapper.__iter__)r   r   r   r   r
   r   r   r   r   r   r   __classcell__r   r   r   r   r   ,  s    r   )N)NN)NN)NF)NN);loggingr_   ro   
contextlibr   typingr   r   r   r   r   r   r	   r
   r!   Ztorch.nn.functionalrr   rs   rA   Z lightning_utilities.core.importsr   r   Ztorch.utils.datar   r   r   Z9lightning_fabric.plugins.environments.cluster_environmentr   Z"lightning_fabric.utilities.importsr   Z$lightning_fabric.utilities.rank_zeror   Z lightning_fabric.utilities.typesr   r/   rM   Ztorch.distributedr   	getLoggerr   r   rH   r   r4   boolrO   rZ   rR   rQ   ZautogradZFunctionre   rt   rv   r   r   r   r   r   r   r   r   r   <module>   s\   (

5 &1 
  *'