Ë
    þÍ:j�  ã                   óx   — d dl mZ d dlZd dlmZ d dlmZ 	 ddedededeeef   d	ed
efd„Zddeded	ed
efd„Z	y)é    )ÚUnionN)ÚTensor)Ú_r2_score_updateÚsum_squared_obsÚsum_obsÚsum_squared_errorÚnum_obsÚsquaredÚreturnc                 óò   — t        j                  |j                  «      j                  }|t        j                  | ||z  |z  z
  |¬«      z  }|st        j
                  |«      }t        j                  |«      S )aÊ  Computes Relative Squared Error.

    Args:
        sum_squared_obs: Sum of square of all observations
        sum_obs: Sum of all observations
        sum_squared_error: Residual sum of squares
        num_obs: Number of predictions or observations
        squared: Returns RRSE value if set to False.

    Example:
        >>> target = torch.tensor([[0.5, 1], [-1, 1], [7, -6]])
        >>> preds = torch.tensor([[0, 2], [-1, 2], [8, -5]])
        >>> # RSE uses the same update function as R2 score.
        >>> sum_squared_obs, sum_obs, rss, num_obs = _r2_score_update(preds, target)
        >>> _relative_squared_error_compute(sum_squared_obs, sum_obs, rss, num_obs, squared=True)
        tensor(0.0632)

    )Úmin)ÚtorchÚfinfoÚdtypeÚepsÚclampÚsqrtÚmean)r   r   r   r	   r
   ÚepsilonÚrses          ú{/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/regression/rse.pyÚ_relative_squared_error_computer      sc   € ô2 �k‰kÐ+×1Ñ1Ó2×6Ñ6€GØ
œeŸk™k¨/¸GÀgÑ<MÐPWÑ<WÑ*WÐ]dÔeÑ
e€CÙÜ�j‰j˜‹oˆÜ�:‰:�c‹?Ðó    ÚpredsÚtargetc                 óD   — t        | |«      \  }}}}t        |||||¬«      S )a"  Computes the relative squared error (RSE).

    .. math:: \text{RSE} = \frac{\sum_i^N(y_i - \hat{y_i})^2}{\sum_i^N(y_i - \overline{y})^2}

    Where :math:`y` is a tensor of target values with mean :math:`\overline{y}`, and
    :math:`\hat{y}` is a tensor of predictions.

    If `preds` and `targets` are 2D tensors, the RSE is averaged over the second dim.

    Args:
        preds: estimated labels
        target: ground truth labels
        squared: returns RRSE value if set to False
    Return:
        Tensor with RSE

    Example:
        >>> from torchmetrics.functional.regression import relative_squared_error
        >>> target = torch.tensor([3, -0.5, 2, 7])
        >>> preds = torch.tensor([2.5, 0.0, 2, 8])
        >>> relative_squared_error(preds, target)
        tensor(0.0514)

    )r
   )r   r   )r   r   r
   r   r   Úrssr	   s          r   Úrelative_squared_errorr   6   s-   € ô2 .>¸eÀVÓ-LÑ*€O�W˜c 7Ü*¨?¸GÀSÈ'Ð[bÔcÐcr   )T)
Útypingr   r   r   Ú%torchmetrics.functional.regression.r2r   ÚintÚboolr   r   © r   r   ú<module>r$      s†   ðõ ã Ý å Bð ñØðàðð ðð �3˜�;Ñð	ð
 ðð óñ@d &ð d°&ð dÀ4ð dÐSYô dr   