a
    dv$                     @   s  d dl 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 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 dFeee ed edddZ!dGeeee e"f ed edddZ#eeeef dddZ$dHeeeee"ee"e"f f  e"ed eee ee df f  ed d!d"Z%dIeee ed$d%d&Z&dJeee e'eee eee ef f d(d)d*Z(dKeeed ed+d,d-Z)dLee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ed5  ed6d7d8Z*dMee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eeeef f d9d:d;Z+dNeed= ed>d?d@Z,dOeeee  ee" eed  edCdDdEZ-dS )P    )OptionalSequenceTupleUnion)Tensor)Literal)spectral_distortion_index)-error_relative_global_dimensionless_synthesis)image_gradients)peak_signal_noise_ratio)relative_average_spectral_error),root_mean_squared_error_using_sliding_window)spectral_angle_mapper).multiscale_structural_similarity_index_measure#structural_similarity_index_measure)total_variation)universal_image_quality_index)_deprecated_root_import_func   elementwise_mean)r   sumnone)predstargetp	reductionreturnc                 C   s   t dd t| |||dS )zWrapper for deprecated import.

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

    r   imager   r   r   r   )r   r   r    r   r/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torchmetrics/functional/image/_deprecated.py_spectral_distortion_index   s    
r!      )r   r   r   N)r   r   ratior   r   c                 C   s   t dd t| |||dS )a1  Wrapper for deprecated import.

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

    r	   r   r   r   r#   r   )r   r	   r$   r   r   r    ._error_relative_global_dimensionless_synthesis*   s    
r%   )imgr   c                 C   s   t dd t| dS )a  Wrapper for deprecated import.

    >>> import torch
    >>> image = torch.arange(0, 1*1*5*5, dtype=torch.float32)
    >>> image = torch.reshape(image, (1, 1, 5, 5))
    >>> dy, dx = _image_gradients(image)
    >>> dy[0, 0, :, :]
    tensor([[5., 5., 5., 5., 5.],
            [5., 5., 5., 5., 5.],
            [5., 5., 5., 5., 5.],
            [5., 5., 5., 5., 5.],
            [0., 0., 0., 0., 0.]])

    r
   r   r&   )r   r
   r'   r   r   r    _image_gradients?   s    
r(   N      $@.)r   r   
data_rangebaser   dimr   c                 C   s   t dd t| |||||dS )zWrapper for deprecated import.

    >>> from torch import tensor
    >>> pred = tensor([[0.0, 1.0], [2.0, 3.0]])
    >>> target = tensor([[3.0, 2.0], [1.0, 0.0]])
    >>> _peak_signal_noise_ratio(pred, target)
    tensor(2.5527)

    r   r   r   r   r*   r+   r   r,   )r   r   r-   r   r   r    _peak_signal_noise_ratioR   s    
r.      )r   r   window_sizer   c                 C   s   t dd t| ||dS )a  Wrapper for deprecated import.

    >>> import torch
    >>> gen = torch.manual_seed(22)
    >>> preds = torch.rand(4, 3, 16, 16, generator=gen)
    >>> target = torch.rand(4, 3, 16, 16, generator=gen)
    >>> _relative_average_spectral_error(preds, target)
    tensor(5114.6641)

    r   r   r   r   r0   )r   r   r1   r   r   r     _relative_average_spectral_errori   s    
r2   F)r   r   r0   return_rmse_mapr   c                 C   s   t dd t| |||dS )a'  Wrapper for deprecated import.

    >>> import torch
    >>> gen = torch.manual_seed(22)
    >>> preds = torch.rand(4, 3, 16, 16, generator=gen)
    >>> target = torch.rand(4, 3, 16, 16, generator=gen)
    >>> _root_mean_squared_error_using_sliding_window(preds, target)
    tensor(0.3999)

    r   r   r   r   r0   r3   )r   r   r4   r   r   r    -_root_mean_squared_error_using_sliding_windowx   s    
r5   )r   r   r   r   c                 C   s   t dd t| ||dS )a  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)
    >>> _spectral_angle_mapper(preds, target)
    tensor(0.5914)

    r   r   r   r   r   )r   r   r6   r   r   r    _spectral_angle_mapper   s    
r7   T      ?   {Gz?Q?gǺ?g48EG?ga4?g??g9EGr?relu)r=   simple)r   r   gaussian_kernelsigmakernel_sizer   r*   k1k2betas	normalizer   c                 C   s(   t dd t| |||||||||	|
dS )a"  Wrapper for deprecated import.

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

    r   r   r   r   r?   r@   rA   r   r*   rB   rC   rD   rE   )r   r   rF   r   r   r    /_multiscale_structural_similarity_index_measure   s    
rG   )r   r   r?   r@   rA   r   r*   rB   rC   return_full_imagereturn_contrast_sensitivityr   c                 C   s(   t dd t| |||||||||	|
dS )zWrapper for deprecated import.

    >>> import torch
    >>> preds = torch.rand([3, 3, 256, 256])
    >>> target = preds * 0.75
    >>> _structural_similarity_index_measure(preds, target)
    tensor(0.9219)

    r   r   r   r   r?   r@   rA   r   r*   rB   rC   rH   rI   )r   r   rJ   r   r   r    $_structural_similarity_index_measure   s    
rK   r   )meanr   r   N)r&   r   r   c                 C   s   t dd t| |dS )zWrapper for deprecated import.

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

    r   r   r&   r   )r   r   rM   r   r   r    _total_variation   s    

rN   r9   r9   r8   r8   )r   r   rA   r@   r   r   c                 C   s   t dd t| ||||dS )zWrapper for deprecated import.

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

    r   r   r   r   rA   r@   r   )r   r   rQ   r   r   r    _universal_image_quality_index   s    
rR   )r   r   )r"   r   )Nr)   r   N)r/   )r/   F)r   )	Tr8   r9   r   Nr:   r;   r<   r=   )	Tr8   r9   r   Nr:   r;   FF)r   )rO   rP   r   ).typingr   r   r   r   Ztorchr   Ztyping_extensionsr   Z&torchmetrics.functional.image.d_lambdar   Z#torchmetrics.functional.image.ergasr	   Z'torchmetrics.functional.image.gradientsr
   Z"torchmetrics.functional.image.psnrr   Z"torchmetrics.functional.image.raser   Z%torchmetrics.functional.image.rmse_swr   Z!torchmetrics.functional.image.samr   Z"torchmetrics.functional.image.ssimr   r   Z torchmetrics.functional.image.tvr   Z!torchmetrics.functional.image.uqir   Ztorchmetrics.utilities.printsr   intr!   floatr%   r(   r.   r2   boolr5   r7   rG   rK   rN   rR   r   r   r   r    <module>   s       
               

*         &   
