Ë
    þÍ:jM  ã                   ó”   — d dl mZ d dlZd dlmZ d dlmZ dedededeeef   fd	„Zdd
edeeef   de	defd„Z
ddedede	dedef
d„Zy)é    )ÚUnionN)ÚTensor)Ú_check_same_shapeÚpredsÚtargetÚnum_outputsÚreturnc                 óÈ   — t        | |«       |dk(  r"| j                  d«      } |j                  d«      }| |z
  }t        j                  ||z  d¬«      }||j                  d   fS )a  Update and returns variables required to compute Mean Squared Error.

    Check for same shape of input tensors.

    Args:
        preds: Predicted tensor
        target: Ground truth tensor
        num_outputs: Number of outputs in multioutput setting

    é   éÿÿÿÿr   )Údim)r   ÚviewÚtorchÚsumÚshape)r   r   r   ÚdiffÚsum_squared_errors        ú{/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/regression/mse.pyÚ_mean_squared_error_updater      sa   € ô �e˜VÔ$Ø�aÒØ—
‘
˜2“ˆØ—‘˜R“ˆØ�6‰>€DÜŸ	™	 $¨¡+°1Ô5ÐØ˜fŸl™l¨1™oÐ-Ð-ó    r   Únum_obsÚsquaredc                 ó@   — |r| |z  S t        j                  | |z  «      S )a  Compute Mean Squared Error.

    Args:
        sum_squared_error: Sum of square of errors over all observations
        num_obs: Number of predictions or observations
        squared: Returns RMSE value if set to False.

    Example:
        >>> preds = torch.tensor([0., 1, 2, 3])
        >>> target = torch.tensor([0., 1, 2, 2])
        >>> sum_squared_error, num_obs = _mean_squared_error_update(preds, target, num_outputs=1)
        >>> _mean_squared_error_compute(sum_squared_error, num_obs)
        tensor(0.2500)

    )r   Úsqrt)r   r   r   s      r   Ú_mean_squared_error_computer   *   s'   € ñ  +2Ð˜wÑ&Ð^´u·z±zÐBSÐV]ÑB]Ó7^Ð^r   c                 ó@   — t        | ||¬«      \  }}t        |||¬«      S )a÷  Compute mean squared error.

    Args:
        preds: estimated labels
        target: ground truth labels
        squared: returns RMSE value if set to False
        num_outputs: Number of outputs in multioutput setting

    Return:
        Tensor with MSE

    Example:
        >>> from torchmetrics.functional.regression import mean_squared_error
        >>> x = torch.tensor([0., 1, 2, 3])
        >>> y = torch.tensor([0., 1, 2, 2])
        >>> mean_squared_error(x, y)
        tensor(0.2500)

    )r   )r   )r   r   )r   r   r   r   r   r   s         r   Úmean_squared_errorr   =   s)   € ô( "<¸EÀ6ÐWbÔ!cÑÐ�wÜ&Ð'8¸'È7ÔSÐSr   )T)Tr   )Útypingr   r   r   Útorchmetrics.utilities.checksr   ÚintÚtupler   Úboolr   r   © r   r   ú<module>r$      s¢   ðõ ã Ý å ;ð. fð .°fð .È3ð .ÐSXÐY_ÐadÐYdÑSeó .ñ(_°6ð _ÀEÈ#ÈvÈ+ÑDVð _Ðaeð _Ðqwó _ñ&T˜fð T¨fð T¸tð TÐY\ð TÐekô Tr   