a
    d                  
   @   s   d dl mZ d dlZd dlmZ d dlmZ d dlmZmZ eee	eeeeeeef dddZ
eeee	ed	d
dZdeee	edddZdS )    )TupleN)Tensor)_uniform_filter)_rmse_sw_compute_rmse_sw_update)predstargetwindow_sizermse_map
target_sumtotal_imagesreturnc                 C   sD   t | ||d||d\}}}|tjt|||d  dd7 }|||fS )a  Calculate the sum of RMSE map values for the batch of examples and update intermediate states.

    Args:
        preds: Deformed image
        target: Ground truth image
        window_size: Sliding window used for RMSE calculation
        rmse_map: Sum of RMSE map values over all examples
        target_sum: target...
        total_images: Total number of images

    Return:
        Intermediate state of RMSE map
        Updated total number of already processed images

    NZrmse_val_sumr
   r      r   )Zdim)r   torchsumr   )r   r   r	   r
   r   r   _ r   k/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torchmetrics/functional/image/rase.py_rase_update   s
     r   )r
   r   r   r	   r   c                 C   sl   t d| |d\}} || }|d}d| tt| d d }t|d }t||| || f S )a  Compute RASE.

    Args:
        rmse_map: Sum of RMSE map values over all examples
        target_sum: target...
        total_images: Total number of images.
        window_size: Sliding window used for rmse calculation

    Return:
        Relative Average Spectral Error (RASE)

    Nr   r   d   r   )r   meanr   sqrtround)r
   r   r   r	   r   Ztarget_meanZrase_mapZ
crop_slider   r   r   _rase_compute1   s    
r      )r   r   r	   r   c                 C   s   t |trt |tr$|dk r$td|jdd }tj||j|jd}tj||j|jd}tjd|jd}t	| |||||\}}}t
||||S )a  Compute Relative Average Spectral Error (RASE) (RelativeAverageSpectralError_).

    Args:
        preds: Deformed image
        target: Ground truth image
        window_size: Sliding window used for rmse calculation

    Return:
        Relative Average Spectral Error (RASE)

    Example:
        >>> from torchmetrics.functional.image import relative_average_spectral_error
        >>> g = torch.manual_seed(22)
        >>> preds = torch.rand(4, 3, 16, 16)
        >>> target = torch.rand(4, 3, 16, 16)
        >>> relative_average_spectral_error(preds, target)
        tensor(5114.6641)

    Raises:
        ValueError: If ``window_size`` is not a positive integer.

       z<Argument `window_size` is expected to be a positive integer.N)dtypedeviceg        )r   )
isinstanceint
ValueErrorshaper   zerosr   r   Ztensorr   r   )r   r   r	   Z	img_shaper
   r   r   r   r   r   relative_average_spectral_errorG   s    r$   )r   )typingr   r   r   Z$torchmetrics.functional.image.helperr   Z%torchmetrics.functional.image.rmse_swr   r   r    r   r   r$   r   r   r   r   <module>   s   