a
    d                     @   s   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m	Z	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 d dlmZ G dd deZdS )    )AnyDictListOptionalUnionN)Tensor)DataParallelModule)Accelerator)CheckpointIO)	Precision)ParallelStrategy)
TBroadcastTReduce)apply_to_collection)ReduceOpc                       s  e Zd ZdZd%ee eeej  ee	 ee
 d fddZeejdddZeddd	d
ZeedddZeddddZd&eeej edddZd'eee eeeef  edddZeeddddZd(eeedddZd)eeedd d!Ze e!dd"d#d$Z"  Z#S )*DataParallelStrategyzImplements data-parallel training in a single process, i.e., the model gets replicated to each device and
    each gets a split of the data.N)acceleratorparallel_devicescheckpoint_io	precisionc                    s   t  j||d ||d d S )N)r   r   Zcluster_environmentr   r   )super__init__)selfr   r   r   r   	__class__ g/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/lightning_fabric/strategies/dp.pyr   !   s    zDataParallelStrategy.__init__)returnc                 C   s   | j d usJ | j d S )Nr   )r   r   r   r   r   root_device0   s    z DataParallelStrategy.root_devicec                 C   s   d S Nr   r   r   r   r   distributed_sampler_kwargs5   s    z/DataParallelStrategy.distributed_sampler_kwargs)moduler   c                 C   s   t || jdS )zMWraps the given model into a :class:`~torch.nn.parallel.DataParallel` module.)r#   Z
device_ids)r   r   r   r#   r   r   r   setup_module9   s    z!DataParallelStrategy.setup_modulec                 C   s   | | j d S r!   )tor    r$   r   r   r   module_to_device=   s    z%DataParallelStrategy.module_to_device)batchdevicer   c                 C   s   |S r!   r   )r   r(   r)   r   r   r   batch_to_device@   s    z$DataParallelStrategy.batch_to_devicemean)
collectiongroup	reduce_opr   c                 C   s   t t ddd}t|t |S )N)tr   c                 S   s   | j }|   |S r!   )Zdtypefloatr+   r&   )r/   Zoriginal_dtyper   r   r   r+   G   s    z-DataParallelStrategy.all_reduce.<locals>.mean)r   r   )r   r,   r-   r.   r+   r   r   r   
all_reduceD   s    zDataParallelStrategy.all_reduce)argskwargsr   c                 O   s   d S r!   r   )r   r2   r3   r   r   r   barrierM   s    zDataParallelStrategy.barrierr   )objsrcr   c                 C   s   |S r!   r   )r   r5   r6   r   r   r   	broadcastP   s    zDataParallelStrategy.broadcastT)decisionallr   c                 C   s   |S r!   r   )r   r8   r9   r   r   r   reduce_boolean_decisionS   s    z,DataParallelStrategy.reduce_boolean_decision)strategy_registryr   c                 C   s   |j d| | jjd d S )NZdp)description)registerr   __name__)clsr;   r   r   r   register_strategiesV   s    z(DataParallelStrategy.register_strategies)NNNN)N)Nr+   )r   )T)$r>   
__module____qualname____doc__r   r
   r   torchr)   r   r   r   propertyr    r"   r	   r   r%   r'   r   r*   r   r   r   strr1   r4   r   intr7   boolr:   classmethodr   r@   __classcell__r   r   r   r   r      s8        	r   )typingr   r   r   r   r   rD   r   Ztorch.nnr   r	   Zlightning_fabric.acceleratorsr
   Z)lightning_fabric.plugins.io.checkpoint_ior   Z"lightning_fabric.plugins.precisionr   Z$lightning_fabric.strategies.parallelr   Z$lightning_fabric.strategies.strategyr   r   Z%lightning_fabric.utilities.apply_funcr   Z&lightning_fabric.utilities.distributedr   r   r   r   r   r   <module>   s   