a
    d                      @   s   d dl mZ d dlmZ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Zd dlmZ d dlmZ d d	lmZ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! erd dl"m#Z# d dl$m%Z% ne& Z%Z#G dd deZ'dS )    )contextmanager)AnyDict	GeneratorListTuple)Tensor)Module)	OptimizerN)_optimizers_to_device)LightningOptimizer)_LightningModuleWrapperBase$_LightningPrecisionModuleWrapperBase)_FAIRSCALE_AVAILABLE_reinit_optimizers_with_oss)DDPSpawnStrategy)	TrainerFn)MisconfigurationException)rank_zero_deprecation)ShardedDataParallel)OSSc                       s   e Zd ZdZdZeedd fddZddd fd	d
ZddddZe	e
e ee	e
e f dddZe
e e
d dddZeedddZeddddZddddZeeddddZ  ZS )DDPSpawnShardedStrategyz1Optimizer sharded training provided by FairScale.Zddp_sharded_spawnN)argskwargsreturnc                    s   t d t j|i | d S )Na  PyTorch Lightning's sharded implementation using FairScale has been deprecated in v1.9.0 and will be removed in v2.0.0. You can try using the `Trainer(strategy='fsdp_native')` instead. The difference is that native FSDP uses PyTorch's implementation and the current strategy uses FairScale's implementation (which was upstreamed to PyTorch). After removal, `strategy='fsdp'` will use the native version by default.)r   super__init__)selfr   r   	__class__ s/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/strategies/sharded_spawn.pyr   ,   s    z DDPSpawnShardedStrategy.__init__zpl.LightningModule)modelr   c                    s   t stdt |S )Nzn`DDPSpawnShardedStrategy` requires `fairscale` to be installed. Install it by running `pip install fairscale`.)r   r   r   connect)r   r"   r   r    r!   r#   6   s
    zDDPSpawnShardedStrategy.connect)r   c                 C   sb   | j d usJ | | j j t| jtjtfs2J | jt	| j| j
d\| _| _
t| j
| j d S )N)r"   
optimizers)lightning_moduleZsetup_optimizerstrainer
isinstancer"   plZLightningModuler   _setup_model_and_optimizersr   r$   r   Zroot_devicer   r    r    r!   configure_ddp>   s    z%DDPSpawnShardedStrategy.configure_ddp)r"   r$   r   c                 C   s(   |  |}t|fd|i| j}||fS )a  Wraps the model and optimizers with fairscale components.

        Return:
            The model wrapped into a :class:`~fairscale.nn.data_parallel.ShardedDataParallel` module
            and a list of optimizer wrapped in :class:~`fairscale.optim.OSS`.
        Zsharded_optimizer)_wrap_optimizersr   Z_ddp_kwargs)r   r"   r$   r    r    r!   r)   H   s    
z3DDPSpawnShardedStrategy._setup_model_and_optimizersr   )r$   r   c                 C   sH   | j s
J | jd ur*| j jjjtjkr*|S dd |D }t|| j| j	S )Nc                 S   s    g | ]}t |tr|jn|qS r    )r'   r   Z
_optimizer).0or    r    r!   
<listcomp>W       z<DDPSpawnShardedStrategy._wrap_optimizers.<locals>.<listcomp>)
r%   r"   r&   statefnr   ZFITTINGr   Zprecision_pluginZ	num_nodes)r   r$   r    r    r!   r,   S   s
    
z(DDPSpawnShardedStrategy._wrap_optimizersc                 c   sH   t | jtr>| j  dV  W d   qD1 s20    Y  ndV  dS )zBlocks syncing gradients behaviour on backwards pass.

        This is useful for skipping sync when accumulating gradients, reducing communication overhead
        Returns: context manager with sync behaviour off
        N)r'   r"   r   Zno_syncr*   r    r    r!   block_backward_syncZ   s    &z+DDPSpawnShardedStrategy.block_backward_sync)closure_lossr   c                 C   s   d S Nr    )r   r4   r    r    r!   pre_backwardg   s    z$DDPSpawnShardedStrategy.pre_backwardc                 C   s   d S r5   r    r*   r    r    r!   post_training_stepj   s    z*DDPSpawnShardedStrategy.post_training_step)strategy_registryr   c                 C   s.   |j d| ddd |j | j| | jj d d S )NZ.ddp_sharded_spawn_find_unused_parameters_falsezADDP Spawn Sharded Strategy with `find_unused_parameters` as FalseF)descriptionZfind_unused_parameters)r9   )registerstrategy_namer   __name__)clsr8   r    r    r!   register_strategiesm   s    z+DDPSpawnShardedStrategy.register_strategies)r<   
__module____qualname____doc__r;   r   r   r#   r+   r	   r   r
   r   r)   r,   r   r   r3   r   r6   r7   classmethodr   r>   __classcell__r    r    r   r!   r   '   s   

"r   )(
contextlibr   typingr   r   r   r   r   Ztorchr   Ztorch.nnr	   Ztorch.optimr
   Zpytorch_lightningr(   Z$lightning_fabric.utilities.optimizerr   Z pytorch_lightning.core.optimizerr   Z pytorch_lightning.overrides.baser   r   Z%pytorch_lightning.overrides.fairscaler   r   Z&pytorch_lightning.strategies.ddp_spawnr   Z pytorch_lightning.trainer.statesr   Z&pytorch_lightning.utilities.exceptionsr   Z%pytorch_lightning.utilities.rank_zeror   Z&fairscale.nn.data_parallel.sharded_ddpr   Zfairscale.optimr   objectr   r    r    r    r!   <module>   s$   