Ë
    þÍ:j7
  ã            	       ór   — d dl Z d dl mZ d dlmZ dededeeef   fd„Z	 dded	ed
edefd„Zdededefd„Zy)é    N)ÚTensor)Ú_check_same_shapeÚpredsÚtargetÚreturnc                 ó    — t        | |«       | |z
  j                  «       j                  «       }|j                  «       j                  «       }||fS )zÕUpdate and returns variables required to compute Weighted Absolute Percentage Error.

    Check for same shape of input tensors.

    Args:
        preds: Predicted tensor
        target: Ground truth tensor

    )r   ÚabsÚsum©r   r   Úsum_abs_errorÚ	sum_scales       ú}/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/regression/wmape.pyÚ/_weighted_mean_absolute_percentage_error_updater      sH   € ô �e˜VÔ$à˜V‘^×(Ñ(Ó*×.Ñ.Ó0€MØ—
‘
“× Ñ Ó"€Ià˜)Ð#Ð#ó    r   r   Úepsilonc                 ó6   — | t        j                  ||¬«      z  S )zãCompute Weighted Absolute Percentage Error.

    Args:
        sum_abs_error: scalar with sum of absolute errors
        sum_scale: scalar with sum of target values
        epsilon: small float to prevent division by zero

    )Úmin)ÚtorchÚclamp)r   r   r   s      r   Ú0_weighted_mean_absolute_percentage_error_computer   *   s   € ð œ5Ÿ;™; y°gÔ>Ñ>Ð>r   c                 ó8   — t        | |«      \  }}t        ||«      S )a½  Compute weighted mean absolute percentage error (`WMAPE`_).

    The output of WMAPE metric is a non-negative floating point, where the optimal value is 0. It is computes as:

    .. math::
        \text{WMAPE} = \frac{\sum_{t=1}^n | y_t - \hat{y}_t | }{\sum_{t=1}^n |y_t| }

    Where :math:`y` is a tensor of target values, and :math:`\hat{y}` is a tensor of predictions.

    Args:
        preds: estimated labels
        target: ground truth labels

    Return:
        Tensor with WMAPE.

    Example:
        >>> from torch import randn
        >>> preds = randn(20,)
        >>> target = randn(20,)
        >>> weighted_mean_absolute_percentage_error(preds, target)
        tensor(1.3967)

    )r   r   r   s       r   Ú'weighted_mean_absolute_percentage_errorr   :   s$   € ô2  OÈuÐV\Ó]Ñ€M�9Ü;¸MÈ9ÓUÐUr   )g-`Àš¡³>)	r   r   Útorchmetrics.utilities.checksr   Útupler   Úfloatr   r   © r   r   ú<module>r      s„   ðó Ý å ;ð$Øð$àð$ð ˆ6�6ˆ>Ñó$ð0 ñ?Øð?àð?ð ð?ð ó	?ð V°6ð VÀ6ð VÈfô Vr   