
    iM                         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        {/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchmetrics/functional/regression/mse.py_mean_squared_error_updater      sa     eV$a

2R6>D		$+15fll1o--    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&^uzzBSV]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)    ( "<E6Wb!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 .SXY_adYdSe .(_6 _E#v+DV _ae _qw _&Tf Tf Tt TY\ Tek Tr   