
    i                         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   fd	Z	d
edeeef   defdZ
dededefdZy)    )UnionN)Tensor)_check_same_shapepredstargetepsilonreturnc                     t        | |       t        j                  | |z
        }|t        j                  t        j                  |      |      z  }t        j                  |      }|j                         }||fS )ac  Update and returns variables required to compute Mean Percentage Error.

    Check for same shape of input tensors.

    Args:
        preds: Predicted tensor
        target: Ground truth tensor
        epsilon: Specifies the lower bound for target values. Any target value below epsilon
            is set to epsilon (avoids ``ZeroDivisionError``).

    )min)r   torchabsclampsumnumel)r   r   r   abs_diffabs_per_errorsum_abs_per_errornum_obss          |/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchmetrics/functional/regression/mape.py&_mean_absolute_percentage_error_updater      sc      eV$yy(Hu{{599V+<'JJM		-0llnGg%%    r   r   c                     | |z  S )aF  Compute Mean Absolute Percentage Error.

    Args:
        sum_abs_per_error: Sum of absolute value of percentage errors over all observations
            ``(percentage error = (target - prediction) / target)``
        num_obs: Number of predictions or observations

    Example:
        >>> target = torch.tensor([1, 10, 1e6])
        >>> preds = torch.tensor([0.9, 15, 1.2e6])
        >>> sum_abs_per_error, num_obs = _mean_absolute_percentage_error_update(preds, target)
        >>> _mean_absolute_percentage_error_compute(sum_abs_per_error, num_obs)
        tensor(0.2667)

     )r   r   s     r   '_mean_absolute_percentage_error_computer   2   s      w&&r   c                 8    t        | |      \  }}t        ||      S )a  Compute mean absolute percentage error.

    Args:
        preds: estimated labels
        target: ground truth labels

    Return:
        Tensor with MAPE

    Note:
        The epsilon value is taken from `scikit-learn's implementation of MAPE`_.

    Example:
        >>> from torchmetrics.functional.regression import mean_absolute_percentage_error
        >>> target = torch.tensor([1, 10, 1e6])
        >>> preds = torch.tensor([0.9, 15, 1.2e6])
        >>> mean_absolute_percentage_error(preds, target)
        tensor(0.2667)

    )r   r   )r   r   r   r   s       r   mean_absolute_percentage_errorr   E   s%    * "Hv!Vw23DgNNr   )g-`>)typingr   r   r   torchmetrics.utilities.checksr   floattupleintr   r   r   r   r   r   <module>r"      s       ; &&& & 63;	&8'v 'PUVY[aVaPb 'gm '&O& O& OV Or   