a
    þdD  ã                   @   sp   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	 d dl
mZ d dlmZ G d	d
„ d
eƒZdS )é    )Úcontextmanager)Ú	GeneratorN)ÚTensor)ÚModule)ÚLiteral)Ú	Precision)Ú_convert_fp_tensorc                   @   sX   e Zd ZU dZdZed ed< eedœdd„Ze	e
d dœd	d
„ƒZeedœdd„ZdS )ÚDoublePrecisionz>Plugin for training with double (``torch.float64``) precision.Z64Ú	precision)ÚmoduleÚreturnc                 C   s   |  ¡ S ©N)Údouble)Úselfr   © r   úr/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/lightning_fabric/plugins/precision/double.pyÚconvert_module   s    zDoublePrecision.convert_module)NNN)r   c                 c   s(   t  ¡ }t  t j¡ dV  t  |¡ dS )zqA context manager to change the default tensor type.

        See: :meth:`torch.set_default_tensor_type`
        N)ÚtorchZget_default_dtypeZset_default_dtypeÚfloat64)r   Zdefault_dtyper   r   r   Úforward_context"   s    zDoublePrecision.forward_context)Údatar   c                 C   s   t |tjƒS r   )r   r   r   )r   r   r   r   r   Úconvert_input-   s    zDoublePrecision.convert_inputN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r
   r   Ú__annotations__r   r   r   r   r   r   r   r   r   r   r   r	      s   

r	   )Ú
contextlibr   Útypingr   r   r   Ztorch.nnr   Ztyping_extensionsr   Z,lightning_fabric.plugins.precision.precisionr   Z(lightning_fabric.plugins.precision.utilsr   r	   r   r   r   r   Ú<module>   s   