a
    dN"                     @   sL  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 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 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G dd deZ G dd deZ!G dd deZ"G dd deZ#G dd de
Z$G dd deZ%G dd deZ&G dd  d eZ'd!S )"    )AnyDictOptionalSequenceTupleUnion)Literal)SpectralDistortionIndex))ErrorRelativeGlobalDimensionlessSynthesis)PeakSignalNoiseRatio)RelativeAverageSpectralError)&RootMeanSquaredErrorUsingSlidingWindow)SpectralAngleMapper)*MultiScaleStructuralSimilarityIndexMeasure StructuralSimilarityIndexMeasure)TotalVariation)UniversalImageQualityIndex)_deprecated_root_import_classc                       s:   e Zd ZdZd	eeef ed edd fddZ	  Z
S )
*_ErrorRelativeGlobalDimensionlessSynthesisa  Wrapper for deprecated import.

    >>> import torch
    >>> preds = torch.rand([16, 1, 16, 16], generator=torch.manual_seed(42))
    >>> target = preds * 0.75
    >>> ergas = _ErrorRelativeGlobalDimensionlessSynthesis()
    >>> torch.round(ergas(preds, target))
    tensor(154.)

       elementwise_meanr   sumnoneNN)ratio	reductionkwargsreturnc                    s&   t dd t jf ||d| d S )Nr
   image)r   r   r   super__init__)selfr   r   r   	__class__ g/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torchmetrics/image/_deprecated.pyr!      s    
z3_ErrorRelativeGlobalDimensionlessSynthesis.__init__)r   r   )__name__
__module____qualname____doc__r   intfloatr   r   r!   __classcell__r%   r%   r#   r&   r      s     
r   c                       sx   e Zd ZdZdeeeee f eeee f e	d e
eeeeef f  eeeedf e	d edd fddZ  ZS )+_MultiScaleStructuralSimilarityIndexMeasurea&  Wrapper for deprecated import.

    >>> import torch
    >>> preds = torch.rand([3, 3, 256, 256], generator=torch.manual_seed(42))
    >>> target = preds * 0.75
    >>> ms_ssim = _MultiScaleStructuralSimilarityIndexMeasure(data_range=1.0)
    >>> ms_ssim(preds, target)
    tensor(0.9627)

    T         ?r   N{Gz?Q?gǺ?g48EG?ga4?g??g9EGr?relur   .)r4   simpleN)gaussian_kernelkernel_sizesigmar   
data_rangek1k2betas	normalizer   r   c
                    s4   t dd t jf |||||||||	d	|
 d S )Nr   r   )	r6   r7   r8   r   r9   r:   r;   r<   r=   r   )r"   r6   r7   r8   r   r9   r:   r;   r<   r=   r   r#   r%   r&   r!   3   s    

z4_MultiScaleStructuralSimilarityIndexMeasure.__init__)	Tr/   r0   r   Nr1   r2   r3   r4   )r'   r(   r)   r*   boolr   r+   r   r,   r   r   r   r   r!   r-   r%   r%   r#   r&   r.   '   s.            
r.   c                
       s^   e Zd ZdZd
eeeeeef f  eed eee	ee	df f  e
dd fdd	Z  ZS )_PeakSignalNoiseRatiozWrapper for deprecated import.

    >>> from torch import tensor
    >>> psnr = _PeakSignalNoiseRatio()
    >>> preds = tensor([[0.0, 1.0], [2.0, 3.0]])
    >>> target = tensor([[3.0, 2.0], [1.0, 0.0]])
    >>> psnr(preds, target)
    tensor(2.5527)

    N      $@r   r   .)r9   baser   dimr   r   c                    s*   t dd t jf ||||d| d S )Nr   r   )r9   rA   r   rB   r   )r"   r9   rA   r   rB   r   r#   r%   r&   r!   [   s    
z_PeakSignalNoiseRatio.__init__)Nr@   r   N)r'   r(   r)   r*   r   r   r,   r   r   r+   r   r!   r-   r%   r%   r#   r&   r?   O   s       r?   c                       s4   e Zd ZdZdeeeef dd fddZ  Z	S )_RelativeAverageSpectralErrora  Wrapper for deprecated import.

    >>> import torch
    >>> g = torch.manual_seed(22)
    >>> preds = torch.rand(4, 3, 16, 16)
    >>> target = torch.rand(4, 3, 16, 16)
    >>> rase = _RelativeAverageSpectralError()
    >>> rase(preds, target)
    tensor(5114.6641)

       Nwindow_sizer   r   c                    s$   t dd t jf d|i| d S )Nr   r   rF   r   r"   rF   r   r#   r%   r&   r!   t   s    
z&_RelativeAverageSpectralError.__init__)rD   
r'   r(   r)   r*   r+   r   strr   r!   r-   r%   r%   r#   r&   rC   g   s    
rC   c                       s4   e Zd ZdZdeeeef dd fddZ  Z	S )'_RootMeanSquaredErrorUsingSlidingWindowa  Wrapper for deprecated import.

    >>> import torch
    >>> g = torch.manual_seed(22)
    >>> preds = torch.rand(4, 3, 16, 16)
    >>> target = torch.rand(4, 3, 16, 16)
    >>> rmse_sw = RootMeanSquaredErrorUsingSlidingWindow()
    >>> rmse_sw(preds, target)
    tensor(0.3999)

    rD   NrE   c                    s$   t dd t jf d|i| d S )Nr   r   rF   r   rG   r#   r%   r&   r!      s    
z0_RootMeanSquaredErrorUsingSlidingWindow.__init__)rD   rH   r%   r%   r#   r&   rJ   }   s    
rJ   c                       s0   e Zd ZdZded edd fddZ  ZS )	_SpectralAngleMappera(  Wrapper for deprecated import.

    >>> import torch
    >>> gen = torch.manual_seed(42)
    >>> preds = torch.rand([16, 3, 16, 16], generator=gen)
    >>> target = torch.rand([16, 3, 16, 16], generator=gen)
    >>> sam = _SpectralAngleMapper()
    >>> sam(preds, target)
    tensor(0.5914)

    r   r   r   r   Nr   r   r   c                    s$   t dd t jf d|i| d S )Nr   r   r   r   r"   r   r   r#   r%   r&   r!      s    
z_SpectralAngleMapper.__init__)r   r'   r(   r)   r*   r   r   r!   r-   r%   r%   r#   r&   rK      s    rK   c                       s2   e Zd ZdZd	eed edd fddZ  ZS )
_SpectralDistortionIndexa  Wrapper for deprecated import.

    >>> import torch
    >>> _ = torch.manual_seed(42)
    >>> preds = torch.rand([16, 3, 16, 16])
    >>> target = torch.rand([16, 3, 16, 16])
    >>> sdi = _SpectralDistortionIndex()
    >>> sdi(preds, target)
    tensor(0.0234)

       r   rL   N)pr   r   r   c                    s&   t dd t jf ||d| d S )Nr	   r   )rR   r   r   )r"   rR   r   r   r#   r%   r&   r!      s    
z!_SpectralDistortionIndex.__init__)rQ   r   )	r'   r(   r)   r*   r+   r   r   r!   r-   r%   r%   r#   r&   rP      s    
rP   c                       sl   e Zd ZdZdeeeee f eeee f e	d
 e
eeeeef f  eeeeedd fddZ  ZS )!_StructuralSimilarityIndexMeasurezWrapper for deprecated import.

    >>> import torch
    >>> preds = torch.rand([3, 3, 256, 256])
    >>> target = preds * 0.75
    >>> ssim = _StructuralSimilarityIndexMeasure(data_range=1.0)
    >>> ssim(preds, target)
    tensor(0.9219)

    Tr0   r/   r   Nr1   r2   Fr   )r6   r8   r7   r   r9   r:   r;   return_full_imagereturn_contrast_sensitivityr   r   c
                    s4   t dd t jf |||||||||	d	|
 d S )Nr   r   )	r6   r8   r7   r   r9   r:   r;   rT   rU   r   )r"   r6   r8   r7   r   r9   r:   r;   rT   rU   r   r#   r%   r&   r!      s    

z*_StructuralSimilarityIndexMeasure.__init__)	Tr0   r/   r   Nr1   r2   FF)r'   r(   r)   r*   r>   r   r,   r   r+   r   r   r   r   r!   r-   r%   r%   r#   r&   rS      s.            rS   c                       s0   e Zd ZdZded edd fddZ  ZS )	_TotalVariationzWrapper for deprecated import.

    >>> import torch
    >>> _ = torch.manual_seed(42)
    >>> tv = _TotalVariation()
    >>> img = torch.rand(5, 3, 28, 28)
    >>> tv(img)
    tensor(7546.8018)

    r   )meanr   r   NNrM   c                    s$   t dd t jf d|i| d S )Nr   r   r   r   rN   r#   r%   r&   r!      s    
z_TotalVariation.__init__)r   rO   r%   r%   r#   r&   rV      s   rV   c                       s<   e Zd ZdZd
ee ee ed edd fdd	Z	  Z
S )_UniversalImageQualityIndexzWrapper for deprecated import.

    >>> import torch
    >>> preds = torch.rand([16, 1, 16, 16])
    >>> target = preds * 0.75
    >>> uqi = _UniversalImageQualityIndex()
    >>> uqi(preds, target)
    tensor(0.9216)

    r/   r/   r0   r0   r   r   N)r7   r8   r   r   r   c                    s(   t dd t jf |||d| d S )Nr   r   )r7   r8   r   r   )r"   r7   r8   r   r   r#   r%   r&   r!     s    
z$_UniversalImageQualityIndex.__init__)rY   rZ   r   )r'   r(   r)   r*   r   r+   r,   r   r   r!   r-   r%   r%   r#   r&   rX      s      rX   N)(typingr   r   r   r   r   r   Ztyping_extensionsr   Ztorchmetrics.image.d_lambdar	   Ztorchmetrics.image.ergasr
   Ztorchmetrics.image.psnrr   Ztorchmetrics.image.raser   Ztorchmetrics.image.rmse_swr   Ztorchmetrics.image.samr   Ztorchmetrics.image.ssimr   r   Ztorchmetrics.image.tvr   Ztorchmetrics.image.uqir   Ztorchmetrics.utilities.printsr   r   r.   r?   rC   rJ   rK   rP   rS   rV   rX   r%   r%   r%   r&   <module>   s*    ((