a
    d                     @   s   d dl mZ 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 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 G dd deZdS )    )contextmanager)AnycastDict	GeneratorOptionalN)Tensor)Module)LBFGS)Literal)_patch_cuda_is_available)	Precision)_convert_fp_tensor)Optimizablec                       s   e Zd ZdZded eeejj	j
 ddddZeed dd	d
ZeedddZeee eedd fddZeeed fddZeeef dddZeeef ddddZejdddZ  ZS )MixedPrecisiona7  Plugin for Automatic Mixed Precision (AMP) training with ``torch.autocast``.

    Args:
        precision: Whether to use ``torch.float16`` (``16``) or ``torch.bfloat16`` (``'bf16'``).
        device: The device for ``torch.autocast``.
        scaler: An optional :class:`torch.cuda.amp.GradScaler` to use.
    N)16   bf16)	precisiondevicescalerreturnc                 C   s   t td t|| _|d u rX| jdkrXt  tjj }W d    n1 sN0    Y  |d urz| jdkrzt	d| d|| _
|| _d S )N)r   r   r   r   z0`precision='bf16'` does not use a scaler, found .)r   r   strr   r   torchcudaamp
GradScaler
ValueErrorr   r   )selfr   r   r    r    v/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/lightning_fabric/plugins/precision/native_amp.py__init__&   s    *zMixedPrecision.__init__)NNN)r   c                 c   s2   |    d V  W d    n1 s$0    Y  d S N)_autocast_context_managerr   r    r    r!   forward_context3   s    
zMixedPrecision.forward_context)datar   c                 C   s"   t jt jd}|| j }t||S )N)r   r   )r   bfloat16float16r   r   )r   r'   Zprecision_to_typeZdst_typer    r    r!   convert_input8   s    
zMixedPrecision.convert_input)tensormodelargskwargsr   c                    s6   | j d ur| j |}t j||g|R i | d S r#   )r   Zscalesuperbackward)r   r+   r,   r-   r.   	__class__r    r!   r0   =   s    
zMixedPrecision.backward)	optimizerr.   r   c                    sR   | j d u rt j|fi |S t|tr0td| j j|fi |}| j   |S )Nz6Native AMP and the LBFGS optimizer are not compatible.)r   r/   optimizer_step
isinstancer
   	TypeErrorstepupdate)r   r3   r.   Zstep_outputr1   r    r!   r4   B   s    


zMixedPrecision.optimizer_stepc                 C   s   | j d ur| j  S i S r#   )r   
state_dictr%   r    r    r!   r9   Q   s    

zMixedPrecision.state_dict)r9   r   c                 C   s   | j d ur| j | d S r#   )r   load_state_dict)r   r9   r    r    r!   r:   V   s    
zMixedPrecision.load_state_dictc                 C   s"   t j| j| jdkrt jnt jdS )Nr   )Zdtype)r   autocastr   r   r(   Zhalfr%   r    r    r!   r$   Z   s    z(MixedPrecision._autocast_context_manager)N)__name__
__module____qualname____doc__r   r   r   r   r   r   r   r"   r   r   r&   r   r*   r	   r   r0   r   r4   r   r9   r:   r;   r$   __classcell__r    r    r1   r!   r      s"   	 r   )
contextlibr   typingr   r   r   r   r   r   r   Ztorch.nnr	   Ztorch.optimr
   Ztyping_extensionsr   Z"lightning_fabric.accelerators.cudar   Z,lightning_fabric.plugins.precision.precisionr   Z(lightning_fabric.plugins.precision.utilsr   Z lightning_fabric.utilities.typesr   r   r    r    r    r!   <module>   s   