
    i'              
           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
edeeef   defdZ	d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      }| j                  r| n| j                         } |j                  r|n|j                         }t	        j
                  t	        j                  | |z
        d      }||j                  d   fS )a  Update and returns variables required to compute Mean Absolute 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is_floating_pointfloattorchsumabsshape)r   r   r   sum_abs_errors       {/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchmetrics/functional/regression/mae.py_mean_absolute_error_updater      s     eV$a

2R,,E%++-E//VV\\^FIIeii7Q?M&,,q/))    r   num_obsc                     | |z  S )a  Compute Mean Absolute Error.

    Args:
        sum_abs_error: Sum of absolute value of errors over all observations
        num_obs: Number of predictions or observations

    Example:
        >>> preds = torch.tensor([0., 1, 2, 3])
        >>> target = torch.tensor([0., 1, 2, 2])
        >>> sum_abs_error, num_obs = _mean_absolute_error_update(preds, target, num_outputs=1)
        >>> _mean_absolute_error_compute(sum_abs_error, num_obs)
        tensor(0.2500)

     )r   r   s     r   _mean_absolute_error_computer   +   s     7""r   c                 <    t        | ||      \  }}t        ||      S )a  Compute mean absolute error.

    Args:
        preds: estimated labels
        target: ground truth labels
        num_outputs: Number of outputs in multioutput setting

    Return:
        Tensor with MAE

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

    )r   )r   r   )r   r   r   r   r   s        r   mean_absolute_errorr   =   s%    & 9T_`M7'w??r   )r   )typingr   r   r   torchmetrics.utilities.checksr   inttupler   r   r   r   r   r   <module>r#      s       ;*v *v *C *TYZ`beZeTf **# #sF{AS #X^ #$@v @v @C @PV @r   