a
    d<                     @   s  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
 d dl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Zd dlmZmZ d d	lmZ d d
lmZ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' d dl(m)Z) d dl*m+Z+ d dl,m-Z- d dl.m/Z/ d dl0m1Z1 d dl2m3Z3 d dl4m5Z5 d dl6m7Z7 d dl8m9Z9 d dl:m;Z;m<Z< d dl=m>Z>m?Z?m@Z@mAZA e BeCZDdZEG dd de3ZFdS )    N)	timedelta)AnyCallableDictListOptionalUnion)Tensor)Module)DistributedDataParallel)Literal)CheckpointIOClusterEnvironment)default_pg_timeout)_distributed_available-_get_default_process_group_backend_for_device_init_dist_connection_sync_ddp_if_availablegroup)_TORCH_GREATER_EQUAL_1_11)_optimizers_to_device)ReduceOp)LightningDistributedModule)$_LightningPrecisionModuleWrapperBase)prepare_for_backward)PrecisionPlugin)_MultiProcessingLauncher)ParallelStrategy)
TBroadcast)	TrainerFn)register_ddp_comm_hook)rank_zero_inforank_zero_only)PredictStepSTEP_OUTPUTTestStepValidationStep)ddp_fork%ddp_fork_find_unused_parameters_falseddp_notebook)ddp_notebook_find_unused_parameters_falsec                       s  e Zd ZdZdZdddddddddedfed eeej	  ee
 ee ee ee ee ee ee ee ed ed fdd	Zeed
ddZejeddddZeed
ddZeej	d
ddZeed
ddZeeeef d
ddZeed
ddZeee d
ddZdd
ddZ dd
 fddZ!ddd d!d"Z"e#e$d#d$d%Z%dd
d&d'Z&dd
d(d)Z'ed
d*d+Z(dd
d,d-Z)dd
d.d/Z*dd
d0d1Z+eee  d
d2d3Z,eedd4d5d6Z-dSe.ee.d8d9d:Z/dd
d;d<Z0e1dd=d>d?Z2dTe1ee ee3e4ef  e1dAdBdCZ5eee6d4dDdEZ7eeee6 d4dFdGZ8eeee6 d4dHdIZ9eee6d4dJdKZ:dd
dLdMZ;e<eddNdOdPZ=dd
 fdQdRZ>  Z?S )UDDPSpawnStrategyzvSpawns processes using the :func:`torch.multiprocessing.spawn` method and joins processes after training
    finishes.	ddp_spawnNspawnzpl.accelerators.Accelerator)r.   forkZ
forkserver)acceleratorparallel_devicescluster_environmentcheckpoint_ioprecision_pluginddp_comm_stateddp_comm_hookddp_comm_wrapperprocess_group_backendtimeoutstart_methodkwargsc                    sP   t  j|||||d d| _|| _|| _|| _|| _d| _|	| _|
| _	|| _
d S )N)r0   r1   r2   r3   r4      r   )super__init__
_num_nodes_ddp_kwargs_ddp_comm_state_ddp_comm_hook_ddp_comm_wrapper_local_rank_process_group_backend_timeout_start_method)selfr0   r1   r2   r3   r4   r5   r6   r7   r8   r9   r:   r;   	__class__ o/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/strategies/ddp_spawn.pyr>   C   s     zDDPSpawnStrategy.__init__)returnc                 C   s   | j S Nr?   rH   rK   rK   rL   	num_nodesc   s    zDDPSpawnStrategy.num_nodes)rQ   rM   c                 C   s
   || _ d S rN   rO   )rH   rQ   rK   rK   rL   rQ   g   s    c                 C   s   | j S rN   )rD   rP   rK   rK   rL   
local_rankl   s    zDDPSpawnStrategy.local_rankc                 C   s   | j d usJ | j | j S rN   )r1   rR   rP   rK   rK   rL   root_devicep   s    zDDPSpawnStrategy.root_devicec                 C   s   | j d urt| j S dS )Nr   )r1   lenrP   rK   rK   rL   num_processesu   s    zDDPSpawnStrategy.num_processesc                 C   s   t | j| j | jd}|S )N)Znum_replicasrank)dictrQ   rU   global_rank)rH   distributed_sampler_kwargsrK   rK   rL   rY   y   s    z+DDPSpawnStrategy.distributed_sampler_kwargsc                 C   s   dS NTrK   rP   rK   rK   rL    _is_single_process_single_device~   s    z1DDPSpawnStrategy._is_single_process_single_devicec                 C   s   | j S rN   )rE   rP   rK   rK   rL   r8      s    z&DDPSpawnStrategy.process_group_backendc                 C   s   t | | jd| _d S )N)r:   )r   rG   Z	_launcherrP   rK   rK   rL   _configure_launcher   s    z$DDPSpawnStrategy._configure_launcherc                    s   |    t   d S rN   )setup_distributedr=   setup_environmentrP   rI   rK   rL   r^      s    z"DDPSpawnStrategy.setup_environmentz
pl.Trainer)trainerrM   c                 C   s   | j d usJ t| j jtjd< | jd us.J | j| |   |jj	}|t
jkrx| jrx| jd ushJ | j| j| _|   |t
jkr|   d S )NZMASTER_PORT)r2   strZ	main_portosenvironr0   setupmodel_to_devicestatefnr    FITTING_layer_syncmodelapplyZsetup_precision_pluginconfigure_ddp)rH   r_   Z
trainer_fnrK   rK   rL   rc      s    

zDDPSpawnStrategy.setup)ri   rM   c                 C   s   t f ||  d| jS )z^Wraps the model into a :class:`~torch.nn.parallel.distributed.DistributedDataParallel` module.)module
device_ids)r   determine_ddp_device_idsr@   )rH   ri   rK   rK   rL   _setup_model   s    zDDPSpawnStrategy._setup_modelc                 C   s\   t | jj d |   | jt_|  | _	| j
d us<J t| j
| j	| j| j| jd d S )Nz: setting up distributed...)r9   )logdetailrJ   __name__set_world_ranksrX   r#   rV   _get_process_group_backendrE   r2   r   Z
world_sizerF   rP   rK   rK   rL   r]      s    
z"DDPSpawnStrategy.setup_distributedc                 C   sL   | j d u rd S | j | j| j | j  | j | j| j  | j  t_	d S rN   )
r2   Zset_global_rankZ	node_rankrU   rR   Zset_world_sizerQ   rX   r#   rV   rP   rK   rK   rL   rs      s
    
z DDPSpawnStrategy.set_world_ranksc                 C   s   | j pt| jS rN   )rE   r   rS   rP   rK   rK   rL   rt      s    z+DDPSpawnStrategy._get_process_group_backendc                 C   s   | j dd| j d< d S )Nfind_unused_parametersT)r@   getrP   rK   rK   rL   pre_configure_ddp   s    z"DDPSpawnStrategy.pre_configure_ddpc                 C   s>   | j jdkr:| jr:t| jts"J t| j| j| j| j	d d S )Ncuda)ri   r5   r6   r7   )
rS   typer[   
isinstanceri   r   r!   rA   rB   rC   rP   rK   rK   rL   _register_ddp_hooks   s    z$DDPSpawnStrategy._register_ddp_hooksc                 C   sf   |    t| jtjtfsJ | t| j| _|   | j	d usFJ | 
| j	j t| j| j d S rN   )rw   rz   ri   plZLightningModuler   ro   r   r{   lightning_moduleZsetup_optimizersr_   r   Z
optimizersrS   rP   rK   rK   rL   rk      s    zDDPSpawnStrategy.configure_ddpc                 C   s   | j jdkrd S | j jgS )Ncpu)rS   ry   indexrP   rK   rK   rL   rn      s    z)DDPSpawnStrategy.determine_ddp_device_ids)argsr;   rM   c                 O   s:   t  s
d S tj dkr,tjj|  d n
tj  d S )NZnccl)rm   )r   torchdistributedget_backendbarrierrn   rH   r   r;   rK   rK   rL   r      s
    zDDPSpawnStrategy.barrierr   )objsrcrM   c                 C   s<   t  s
|S |g}| j|kr d g}tjj||tjd |d S )Nr   r   )r   rX   r   r   Zbroadcast_object_list_groupZWORLD)rH   r   r   rK   rK   rL   	broadcast   s    
zDDPSpawnStrategy.broadcastc                 C   s:   | j jdkrtj| j  | jd us(J | j| j  d S )Nrx   )rS   ry   r   rx   Z
set_deviceri   torP   rK   rK   rL   rd      s    z DDPSpawnStrategy.model_to_device)closure_lossrM   c                 C   s6   t | jtsdS | jdusJ | jjs2t| j| dS )z.Run before precision plugin executes backward.N)rz   ri   r   r}   automatic_optimizationr   )rH   r   rK   rK   rL   pre_backward   s
    zDDPSpawnStrategy.pre_backwardmean)tensorr   	reduce_oprM   c                 C   s   t |trt|||d}|S )a  Reduces a tensor from several distributed processes to one aggregated tensor.

        Args:
            tensor: the tensor to sync and reduce
            group: the process group to gather results from. Defaults to all processes (world)
            reduce_op: the reduction operation. Defaults to 'mean'/'avg'.
                Can also be a string 'sum' to calculate the sum during reduction.

        Return:
            reduced value, except when the input was not a tensor the output remains is unchanged
        )r   )rz   r	   r   )rH   r   r   r   rK   rK   rL   reduce  s    
zDDPSpawnStrategy.reducec                 O   sL   | j d usJ | j   | j |i |W  d    S 1 s>0    Y  d S rN   )ri   r4   Ztrain_step_contextr   rK   rK   rL   training_step  s    zDDPSpawnStrategy.training_stepc                 O   s   | j   | jd usJ | jd us(J | jjjjtjkrX| j|i |W  d    S t	| jt
shJ | jj|i |W  d    S W d    n1 s0    Y  d S rN   )r4   Zval_step_contextr}   ri   r_   re   rf   r    rg   rz   r'   validation_stepr   rK   rK   rL   r     s    z DDPSpawnStrategy.validation_stepc                 O   sP   | j  2 t| jtsJ | jj|i |W  d    S 1 sB0    Y  d S rN   )r4   Ztest_step_contextrz   ri   r&   	test_stepr   rK   rK   rL   r   &  s    zDDPSpawnStrategy.test_stepc                 O   sP   | j  2 t| jtsJ | jj|i |W  d    S 1 sB0    Y  d S rN   )r4   Zpredict_step_contextrz   ri   r$   predict_stepr   rK   rK   rL   r   +  s    zDDPSpawnStrategy.predict_stepc                 C   s0   | j d usJ | j js,| jd us$J d| j_d S rZ   )r}   r   ri   Zrequire_backward_grad_syncrP   rK   rK   rL   post_training_step0  s    z#DDPSpawnStrategy.post_training_step)strategy_registryrM   c                 C   s^   d}|D ]"\}}|j || d| d|d qd}|D ]$\}}|j || d| dd|d q4d S )	N))r-   r.   )r(   r/   )r*   r/   z"DDP strategy with `start_method` '')descriptionr:   ))Z&ddp_spawn_find_unused_parameters_falser.   )r)   r/   )r+   r/   zHDDP strategy with `find_unused_parameters` as False and `start_method` 'F)r   ru   r:   )register)clsr   entriesnamer:   rK   rK   rL   register_strategies6  s"    

z$DDPSpawnStrategy.register_strategiesc                    s   t | jj d | j}t| jtr\trV| jj	sV| j
 drVtd| jj d || _|d ur|jd ur|jjjtjkr| jr| jd usJ | j| j| _t   d S )Nz: tearing down strategyZcan_set_static_graphzyYour model can run with static graph optimizations. For future training runs, we suggest you pass `Trainer(..., strategy=z%(static_graph=True))` to enable them.)rp   rq   rJ   rr   r}   rz   ri   r   r   Zstatic_graphZ_get_ddp_logging_datarv   r"   Z_trainerre   rf   r    rg   rh   revertr=   teardown)rH   Z	pl_modulerI   rK   rL   r   S  s4    zDDPSpawnStrategy.teardown)r   )Nr   )@rr   
__module____qualname____doc__Zstrategy_namer   r   r   r   Zdevicer   r   r   objectr   r`   r   r   r   r>   propertyintrQ   setterrR   rS   rU   r   rY   boolr[   r8   r\   r^   rc   r
   r   ro   r]   rs   rt   rw   r{   rk   rn   r   r   r   rd   r	   r   r   r   r   r%   r   r   r   r   r   classmethodr   r   __classcell__rK   rK   rI   rL   r,   =   s    		 r,   )Gloggingra   datetimer   typingr   r   r   r   r   r   r   Ztorch.distributedr	   Ztorch.nnr
   Ztorch.nn.parallel.distributedr   Ztyping_extensionsr   Zpytorch_lightningr|   Zlightning_fabric.pluginsr   r   Z5lightning_fabric.plugins.collectives.torch_collectiver   Z&lightning_fabric.utilities.distributedr   r   r   r   r   r   Z"lightning_fabric.utilities.importsr   Z$lightning_fabric.utilities.optimizerr   Z lightning_fabric.utilities.typesr   Zpytorch_lightning.overridesr   Z pytorch_lightning.overrides.baser   Z'pytorch_lightning.overrides.distributedr   Z#pytorch_lightning.plugins.precisionr   Z6pytorch_lightning.strategies.launchers.multiprocessingr   Z%pytorch_lightning.strategies.parallelr   Z%pytorch_lightning.strategies.strategyr   Z pytorch_lightning.trainer.statesr    Z'pytorch_lightning.utilities.distributedr!   Z%pytorch_lightning.utilities.rank_zeror"   r#   Z!pytorch_lightning.utilities.typesr$   r%   r&   r'   	getLoggerrr   rp   Z_DDP_FORK_ALIASESr,   rK   rK   rK   rL   <module>   s>    
