
    i                         d dl Z d dl mZ d dlmZmZ d dlmZ dedededed	ed
edeeeef   fdZ	ded	ed
ededef
dZ
ddedededefdZy)    N)Tensor)_rmse_sw_compute_rmse_sw_update)_uniform_filterpredstargetwindow_sizermse_map
target_sumtotal_imagesreturnc                     t        | ||d||      \  }}}|t        j                  t        ||      |dz  z  d      z  }|||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rmse_val_sumr
   r      r   )dim)r   torchsumr   )r   r   r	   r
   r   r   _s          w/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchmetrics/functional/image/rase.py_rase_updater      sY    $ !0v{Wc!Ax %))OFK@KQRNSYZ[[JZ--    c                    t        d| |      \  }} ||z  }|j                  d      }d|z  t        j                  t        j                  | dz  d            z  }t	        |dz        }t        j                  ||| || 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   target_meanrase_map
crop_slides           r   _rase_computer!   0   s     #xVbcKAx|+K""1%K[ 5::ejj1a.H#IIH{Q'J::hz:+5z:+7MMNOOr   c                    t        |t              rt        |t              r|dk  rt        d      |j                  dd }t	        j
                  ||j                  |j                        }t	        j
                  ||j                  |j                        }t	        j                  d|j                        }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 torch import rand
        >>> from torchmetrics.functional.image import relative_average_spectral_error
        >>> preds = rand(4, 3, 16, 16)
        >>> target = rand(4, 3, 16, 16)
        >>> relative_average_spectral_error(preds, target)
        tensor(5326.40...)

    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%   tensorr   r!   )r   r   r	   	img_shaper
   r   r   s          r   relative_average_spectral_errorr-   F   s    . k3'J{C,H[[\_WXXQR I{{9FLLOHYfll6==QJ<<FMM:L)5eV[RZ\fht)u&Hj,:|[IIr   )   )r   r   %torchmetrics.functional.image.rmse_swr   r   #torchmetrics.functional.image.utilsr   r'   tupler   r!   r-    r   r   <module>r3      s       S ?..!.03.?E.SY.io.
666!".2PF P Pf P[^ Pci P, J6  J6  JPS  J\b  Jr   