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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 G dd deZG dd deZdS )    )contextmanager)Anycast	GeneratorListTupleN)apply_to_collection)FloatTensorTensor)	Optimizer)Literal)_convert_fp_tensor)$_LightningPrecisionModuleWrapperBase)PrecisionPluginc                   @   s~   e Zd ZdZeeedddZeeedddZeeeddd	Zeeedd
dZ	eeedddZ
eeedddZdS )LightningDoublePrecisionModulezLightningModule wrapper which converts incoming floating point data in ``*_step`` and ``forward`` to double
    (``torch.float64``) precision.

    Args:
        pl_module: the model to wrap
    )
collectionreturnc                 C   s   t | tttjdS )N)functionZdst_type)r   r
   r   torchdouble)r    r   s/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/plugins/precision/double.py_move_float_tensors_to_double&   s    z<LightningDoublePrecisionModule._move_float_tensors_to_double)argskwargsr   c                 O   s   | j jt|i t|S N)moduletraining_stepr   r   selfr   r   r   r   r   r   *   s
    z,LightningDoublePrecisionModule.training_stepc                 O   s   | j jt|i t|S r   )r   validation_stepr   r   r   r   r   r   r    0   s
    z.LightningDoublePrecisionModule.validation_stepc                 O   s   | j jt|i t|S r   )r   	test_stepr   r   r   r   r   r   r!   6   s
    z(LightningDoublePrecisionModule.test_stepc                 O   s   | j jt|i t|S r   )r   predict_stepr   r   r   r   r   r   r"   <   s
    z+LightningDoublePrecisionModule.predict_stepc                 O   s   | j t|i t|S r   )r   r   r   r   r   r   r   forwardB   s
    z&LightningDoublePrecisionModule.forwardN)__name__
__module____qualname____doc__staticmethodr   r   r   r    r!   r"   r#   r   r   r   r   r      s   r   c                       sr   e Zd ZU dZdZed ed< eje	e
 e	e eeje	d e	e f d fddZeed d	d
dZ  ZS )DoublePrecisionPluginz>Plugin for training with double (``torch.float64``) precision.Z64	precisionr   )model
optimizerslr_schedulersr   c                    s(   t tj| }t|}t |||S )zConverts the model to double precision and wraps it in a ``LightningDoublePrecisionModule`` to convert
        incoming floating point data to double (``torch.float64``) precision.

        Does not alter `optimizers` or `lr_schedulers`.
        )r   plZLightningModuler   r   superconnect)r   r+   r,   r-   	__class__r   r   r0   N   s    zDoublePrecisionPlugin.connect)NNN)r   c                 c   s    t t j dV  t t dS )zqA context manager to change the default tensor type.

        See: :meth:`torch.set_default_tensor_type`
        N)r   Zset_default_tensor_typeZDoubleTensorr	   )r   r   r   r   forward_context[   s    z%DoublePrecisionPlugin.forward_context)r$   r%   r&   r'   r*   r   __annotations__nnModuler   r   r   r   r0   r   r   r3   __classcell__r   r   r1   r   r)   I   s   
r)   )
contextlibr   typingr   r   r   r   r   r   Ztorch.nnr5   Z#lightning_utilities.core.apply_funcr   r	   r
   Ztorch.optimr   Ztyping_extensionsr   Zpytorch_lightningr.   Z(lightning_fabric.plugins.precision.utilsr   Z pytorch_lightning.overrides.baser   Z4pytorch_lightning.plugins.precision.precision_pluginr   r   r)   r   r   r   r   <module>   s   +