Ë
    þÍ:j'  ã            
       óŠ   — 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       ú{/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/regression/mae.pyÚ_mean_absolute_error_updater      sŠ   € ô �e˜VÔ$Ø�aÒØ—
‘
˜2“ˆØ—‘˜R“ˆØ×,Ò,‰E°%·+±+³-€EØ×/Ò/‰V°V·\±\³^€FÜ—I‘IœeŸi™i¨°©Ó7¸QÔ?€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_Ô`Ñ€M�7Ü'¨°wÓ?Ð?r   )r   )Útypingr   r   r   Útorchmetrics.utilities.checksr   ÚintÚtupler   r   r   r   r   r   ú<module>r#      sŒ   ðõ ã Ý å ;ð* vð *°vð *ÈCð *ÐTYÐZ`ÐbeÐZeÑTfó *ð*#°ð #ÀÀsÈFÀ{ÑASð #ÐX^ó #ñ$@˜vð @¨vð @ÀCð @ÐPVô @r   