a
    dr                     @   s  d dl mZ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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! edZ"ere"rd dl#Z#e! Z$ed Z%ed Z&ee%e&f Z'G dd deZ(dS )    )AnyCallablecastOptionalTYPE_CHECKINGUnion)RequirementCache)Tensor)LBFGS	Optimizer)get_argsLiteralN)	Steppable)_APEX_AVAILABLE)PrecisionPlugin)GradClipAlgorithmType)MisconfigurationException)is_overridden)rank_zero_deprecationWarningCache	deepspeed)       )3216bf16c                   @   s   e Zd ZdZded ee ee ddddZedee	 ee
 eeddd	d
Ze	de
eg ef eedddZdejfeee
ef eddddZdddddZdS )DeepSpeedPrecisionPluginzPrecision plugin for DeepSpeed integration.

    Args:
        precision: Full precision (32), half precision (16) or bfloat16 precision (bf16).
    Raises:
        ValueError:
            If unsupported ``precision`` is provided.
    N)r   r   r   r   r   )	precisionamp_type	amp_levelreturnc                 C   s   |dkr&t d tstd|p"d}n$|d urJtdt| j d|d|d u rXd}nt d	t| j d
|d tttt }||vrtd|d| dt	tt
|| _|| _|| _d S )NZapexzThe NVIDIA/apex AMP implementation has been deprecated upstream. Consequently, its integration inside PyTorch Lightning has been deprecated in v1.9.0. Support for using it through the DeepSpeed implementation will be removed in v2.0.0.zxYou have asked for Apex AMP but `apex` is not installed. Install `apex` using this guide: https://github.com/NVIDIA/apexZO2`z(amp_level=z*)` is only relevant when using NVIDIA/apexZnativez	Passing `z
(amp_type=za)` been deprecated in v1.9.0 and will be removed in v2.0.0. This argument is no longer necessary.z)`Trainer(strategy='deepspeed', precision=z1)` is not supported. `precision` must be one of: .)r   r   r   
ValueErrortype__name__r   _PRECISION_INPUT_STR_PRECISION_INPUT_INTr   strr   r   r   )selfr   r   r   Zsupported_precision r*   v/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/plugins/precision/deepspeed.py__init__3   s8    
z!DeepSpeedPrecisionPlugin.__init__zpl.LightningModule)tensormodel	optimizeroptimizer_idxargskwargsr    c                 O   s8   t d|rtd |jj}|j|g|R i | dS )a  Performs back-propagation using DeepSpeed's engine.

        Args:
            tensor: the loss tensor
            model: the model to be optimized
            optimizer: ignored for DeepSpeed
            optimizer_idx: ignored for DeepSpeed
            \*args: additional positional arguments for the :meth:`deepspeed.DeepSpeedEngine.backward` call
            \**kwargs: additional keyword arguments for the :meth:`deepspeed.DeepSpeedEngine.backward` call
        backwardzYou have overridden the `LightningModule.backward` hook but it will be ignored since DeepSpeed handles the backward logic internally.N)r   warning_cachewarntrainerr.   r3   )r)   r-   r.   r/   r0   r1   r2   deepspeed_enginer*   r*   r+   r3   ]   s    
z!DeepSpeedPrecisionPlugin.backward)r/   r.   r0   closurer2   r    c           	      K   s`   t |trtd| d| }| ||| |d u }|jrH|rHtd|jj}|jf i |S )Nz@DeepSpeed and the LBFGS optimizer are not compatible (optimizer z).z_Skipping backward by returning `None` from your `training_step` is not supported by `DeepSpeed`)
isinstancer
   r   Z_after_closureZautomatic_optimizationr6   r.   step)	r)   r/   r.   r0   r8   r2   Zclosure_resultZskipped_backwardr7   r*   r*   r+   optimizer_stepx   s    


z'DeepSpeedPrecisionPlugin.optimizer_stepg        )r/   clip_valgradient_clip_algorithmr    c                 C   s   dS )z/DeepSpeed handles gradient clipping internally.Nr*   )r)   r/   r<   r=   r*   r*   r+   clip_gradients   s    z'DeepSpeedPrecisionPlugin.clip_gradientsz
pl.Trainer)r6   r    c                 C   s&   |j dkrd S td|j d d S )Nz!You set `Trainer(track_grad_norm=zH)' but this is not supported for DeepSpeed. The setting will be ignored.)Ztrack_grad_normr4   r5   )r)   r6   r*   r*   r+   _track_grad_norm   s
    
z)DeepSpeedPrecisionPlugin._track_grad_norm)NN)r%   
__module____qualname____doc__r   r   r(   r,   r	   r   intr   r3   r   r;   r   ZNORMr   r   floatr>   r@   r*   r*   r*   r+   r   )   sB     ,

r   ))typingr   r   r   r   r   r   Z lightning_utilities.core.importsr   Ztorchr	   Ztorch.optimr
   r   Ztyping_extensionsr   r   Zpytorch_lightningplZ lightning_fabric.utilities.typesr   Z,pytorch_lightning.plugins.precision.apex_ampr   Z4pytorch_lightning.plugins.precision.precision_pluginr   Zpytorch_lightning.utilitiesr   Z&pytorch_lightning.utilities.exceptionsr   Z)pytorch_lightning.utilities.model_helpersr   Z%pytorch_lightning.utilities.rank_zeror   r   Z_DEEPSPEED_AVAILABLEr   r4   r'   r&   Z_PRECISION_INPUTr   r*   r*   r*   r+   <module>   s(    