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	 d dl
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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)apply_to_collection)Tensor)DataParallelModule)CheckpointIO)ReduceOp)$_LightningPrecisionModuleWrapperBase)LightningParallelModule)PrecisionPlugin)ParallelStrategy)
TBroadcastTReduce)is_overridden)STEP_OUTPUTc                       s  e Zd ZdZdZd<ed eeej  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ddd fddZd=eeej eedddZeedddZd>eee eeeef  edddZeejdd d!Zd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e$d$d.d/Z%eeee$ d$d0d1Z&eeee$ d$d2d3Z'eee$d$d4d5Z(e$e$d6d7d8Z)e*e+dd9d:d;Z,  Z-S )A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.ZdpNzpl.accelerators.Accelerator)acceleratorparallel_devicescheckpoint_ioprecision_pluginc                    s   t  j||d ||d d S )N)r   r   Zcluster_environmentr   r   )super__init__)selfr   r   r   r   	__class__ h/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/strategies/dp.pyr   '   s    zDataParallelStrategy.__init__)returnc                 C   s   dS Nr   r   r   r   r   r    global_rank6   s    z DataParallelStrategy.global_rankc                 C   s   dS r"   r   r#   r   r   r    
local_rank:   s    zDataParallelStrategy.local_rankc                 C   s   dS r"   r   r#   r   r   r    	node_rank>   s    zDataParallelStrategy.node_rankc                 C   s   dS )N   r   r#   r   r   r    
world_sizeB   s    zDataParallelStrategy.world_sizez
pl.Trainer)trainerr!   c                    s@   |    t| jtjtfsJ | t| j| _t 	| d S N)
model_to_device
isinstancemodelplZLightningModuler   _setup_modelr   r   setup)r   r)   r   r   r    r0   F   s    zDataParallelStrategy.setupr   )batchdevicedataloader_idxr!   c                 C   s   |S )a2  Moves the batch to the correct device.

        The input and the output is the same type.

        Args:
            batch: The batch of samples to move to the correct device
            device: The target device
            dataloader_idx: The index of the dataloader to which the batch belongs.
        r   )r   r1   r2   r3   r   r   r    batch_to_deviceM   s    z$DataParallelStrategy.batch_to_device)r-   r!   c                 C   s   t || jdS )zMWraps the given model into a :class:`~torch.nn.parallel.DataParallel` module.)moduleZ
device_ids)r	   r   )r   r-   r   r   r    r/   Z   s    z!DataParallelStrategy._setup_modelmean)
collectiongroup	reduce_opr!   c                 C   s   t t ddd}t|t |S )as  Reduces a collection of tensors from all processes. It can be applied to just a single tensor.

        Args:
            collection: The collection of tensors to sync and reduce.
            group: ignored for DP
            reduce_op: ignored for DP
        Return:
            Reduced tensor values or the same value if it was not or did not contain a tensor.
        )tr!   c                 S   s   | j }|   |S r*   )Zdtypefloatr6   to)r:   Zoriginal_dtyper   r   r    r6   k   s    z)DataParallelStrategy.reduce.<locals>.mean)r   r   )r   r7   r8   r9   r6   r   r   r    reduce^   s    zDataParallelStrategy.reducec                 C   s   | j d usJ | j d S r"   )r   r#   r   r   r    root_deviceq   s    z DataParallelStrategy.root_devicec                 C   s    | j d usJ | j | j d S r*   )r-   r<   r>   r#   r   r   r    r+   v   s    z$DataParallelStrategy.model_to_device)argskwargsr!   c                 O   s   d S r*   r   r   r?   r@   r   r   r    barrierz   s    zDataParallelStrategy.barrier)objsrcr!   c                 C   s   |S r*   r   )r   rC   rD   r   r   r    	broadcast}   s    zDataParallelStrategy.broadcastT)decisionallr!   c                 C   s   |S r*   r   )r   rF   rG   r   r   r    reduce_boolean_decision   s    z,DataParallelStrategy.reduce_boolean_decisionc                 O   sL   | j  . | jd usJ | j|i |W  d    S 1 s>0    Y  d S r*   )r   Ztrain_step_contextr-   rA   r   r   r    training_step   s    z"DataParallelStrategy.training_stepc                 O   sL   | j  . | jd usJ | j|i |W  d    S 1 s>0    Y  d S r*   )r   Zval_step_contextr-   rA   r   r   r    validation_step   s    z$DataParallelStrategy.validation_stepc                 O   sL   | j  . | jd usJ | j|i |W  d    S 1 s>0    Y  d S r*   )r   Ztest_step_contextr-   rA   r   r   r    	test_step   s    zDataParallelStrategy.test_stepc                 O   sL   | j  . | jd usJ | j|i |W  d    S 1 s>0    Y  d S r*   )r   Zpredict_step_contextr-   rA   r   r   r    predict_step   s    z!DataParallelStrategy.predict_step)outputr!   c                 C   sN   t d| jr|S t|tr6d|v r6| |d |d< nt|trJ| |}|S )Ntraining_step_endZloss)r   Zlightning_moduler,   dictr=   r   )r   rM   r   r   r    rN      s    

z&DataParallelStrategy.training_step_end)strategy_registryr!   c                 C   s   |j | j| | jj d d S )N)description)registerstrategy_namer   __name__)clsrP   r   r   r    register_strategies   s
    z(DataParallelStrategy.register_strategies)NNNN)Nr   )Nr6   )r   )T).rT   
__module____qualname____doc__rS   r   r   torchr2   r   r   r   propertyintr$   r%   r&   r(   r0   r   r4   r
   r	   r/   r   r   r   strr=   r>   r+   rB   r   rE   boolrH   r   rI   rJ   rK   rL   rN   classmethodr   rV   __classcell__r   r   r   r    r   !   sR        r   )#typingr   r   r   r   r   rZ   Z#lightning_utilities.core.apply_funcr   r   Ztorch.nnr	   r
   Zpytorch_lightningr.   Zlightning_fabric.pluginsr   Z&lightning_fabric.utilities.distributedr   Z pytorch_lightning.overrides.baser   Z)pytorch_lightning.overrides.data_parallelr   Z#pytorch_lightning.plugins.precisionr   Z%pytorch_lightning.strategies.parallelr   Z%pytorch_lightning.strategies.strategyr   r   Z)pytorch_lightning.utilities.model_helpersr   Z!pytorch_lightning.utilities.typesr   r   r   r   r   r    <module>   s   