Ë
    þÍ:jx3  ã                   ó  — d dl mZmZmZ d dlZd dlmZ d dlmZ d dlm	Z	m
Z
mZ d dlmZ deded	efd
„Zdeeee   f   deeee   f   d	efd„Z	 	 	 ddededeeeee   ef      dee   ded	efd„Z	 	 	 	 ddedededeeeee   ef      deed      dee   ded	efd„Z	 	 	 ddedededeeeee   ef      dee   ded	efd„Z	 	 	 	 	 	 ddededed   deeeee   ef      dee   dee   deed      dee   ded	eeee   f   fd„Zy)é    )ÚListÚOptionalÚUnionN)ÚTensor)ÚLiteral)Ú
binary_rocÚmulticlass_rocÚmultilabel_roc)ÚClassificationTaskÚfprÚtprÚreturnc                 ó„   — | d|z
  z
  }t        j                  t        j                  |«      «      }| |   d||   z
  z   dz  S )z>Compute Equal Error Rate (EER) for binary classification task.é   é   )ÚtorchÚargminÚabs)r   r   ÚdiffÚidxs       ú/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/classification/eer.pyÚ_binary_eer_computer      sA   € à�!�c‘'‰?€DÜ
�,‰,”u—y‘y “Ó
'€CØ�‰H˜˜C ™H™Ñ%¨Ñ*Ð*ó    c           
      óü   — t        | t        «      r+t        |t        «      r| j                  dk(  rt        | |«      S t	        j
                  t        | |«      D ��cg c]  \  }}t        ||«      ‘Œ c}}«      S c c}}w )zCompute Equal Error Rate (EER).r   )Ú
isinstancer   Úndimr   r   ÚstackÚzip)r   r   ÚfÚts       r   Ú_eer_computer!   #   s]   € ô
 �#”vÔ¤:¨c´6Ô#:¸s¿x¹xÈ1º}Ü" 3¨Ó,Ð,Ü�;‰;¼cÀ#Às»m×L±d°a¸Ô+¨A¨qÕ1ÓLÓMÐMùÓLs   ÁA8
ÚpredsÚtargetÚ
thresholdsÚignore_indexÚvalidate_argsc                 ó@   — t        | ||||«      \  }}}t        ||«      S )a  Compute Equal Error Rate (EER) for binary classification task.

    .. math::
        \text{EER} = \frac{\text{FAR} + \text{FRR}}{2}, \text{where} \min_t abs(FAR_t-FRR_t)

    The Equal Error Rate (EER) is the point where the False Positive Rate (FPR) and True Positive Rate (TPR) are
    equal, or in practise minimized. A lower EER value signifies higher system accuracy.

    Args:
        preds: Tensor with predictions
        target: Tensor with true labels
        thresholds:
            Can be one of:

            - If set to `None`, will use a non-binned approach where thresholds are dynamically calculated from
              all the data. Most accurate but also most memory consuming approach.
            - If set to an `int` (larger than 1), will use that number of thresholds linearly spaced from
              0 to 1 as bins for the calculation.
            - If set to an `list` of floats, will use the indicated thresholds in the list as bins for the calculation
            - If set to an 1d `tensor` of floats, will use the indicated thresholds in the tensor as
              bins for the calculation.

        ignore_index:
            Specifies a target value that is ignored and does not contribute to the metric calculation
        validate_args: bool indicating if input arguments and tensors should be validated for correctness.
            Set to ``False`` for faster computations

    Returns:
        A single scalar with the eer score

    Example:
        >>> from torchmetrics.functional.classification import binary_eer
        >>> preds = torch.tensor([0, 0.5, 0.7, 0.8])
        >>> target = torch.tensor([0, 1, 1, 0])
        >>> binary_eer(preds, target, thresholds=None)
        tensor(0.5000)
        >>> binary_eer(preds, target, thresholds=5)
        tensor(0.7500)

    )r   r!   )r"   r#   r$   r%   r&   r   r   Ú_s           r   Ú
binary_eerr)   -   s*   € ô^ ˜U F¨J¸ÀmÓT�K€CˆˆaÜ˜˜SÓ!Ð!r   Únum_classesÚaverage)ÚmicroÚmacroc           	      óD   — t        | ||||||«      \  }}}	t        ||«      S )a  Compute Equal Error Rate (EER) for multiclass classification task.

    .. math::
        \text{EER} = \frac{\text{FAR} + (1 - \text{FRR})}{2}, \text{where} \min_t abs(FAR_t-FRR_t)

    The Equal Error Rate (EER) is the point where the False Positive Rate (FPR) and True Positive Rate (TPR) are
    equal, or in practise minimized. A lower EER value signifies higher system accuracy.

    Args:
        preds: Tensor with predictions
        target: Tensor with true labels
        num_classes: Integer specifying the number of classes
        thresholds:
            Can be one of:

            - If set to `None`, will use a non-binned approach where thresholds are dynamically calculated from
              all the data. Most accurate but also most memory consuming approach.
            - If set to an `int` (larger than 1), will use that number of thresholds linearly spaced from
              0 to 1 as bins for the calculation.
            - If set to an `list` of floats, will use the indicated thresholds in the list as bins for the calculation
            - If set to an 1d `tensor` of floats, will use the indicated thresholds in the tensor as
              bins for the calculation.
        average:
            If aggregation of should be applied. The aggregation is applied to underlying ROC curves.
            By default, eer is not aggregated and a score for each class is returned. If `average` is set to ``"micro"``
            , the metric will aggregate the curves by one hot encoding the targets and flattening the predictions,
            considering all classes jointly as a binary problem. If `average` is set to ``"macro"``, the metric will
            aggregate the curves by first interpolating the curves from each class at a combined set of thresholds and
            then average over the classwise interpolated curves. See `averaging curve objects`_ for more info on the
            different averaging methods.
        ignore_index:
            Specifies a target value that is ignored and does not contribute to the metric calculation
        validate_args: bool indicating if input arguments and tensors should be validated for correctness.
            Set to ``False`` for faster computations.

    Returns:
        If `average=None|"none"` then a 1d tensor of shape (n_classes, ) will be returned with eer score per class.
        If `average="macro"|"micro"` then a single scalar is returned.


    Example:
        >>> from torchmetrics.functional.classification import multiclass_eer
        >>> preds = torch.tensor([[0.75, 0.05, 0.05, 0.05, 0.05],
        ...                       [0.05, 0.75, 0.05, 0.05, 0.05],
        ...                       [0.05, 0.05, 0.75, 0.05, 0.05],
        ...                       [0.05, 0.05, 0.05, 0.75, 0.05]])
        >>> target = torch.tensor([0, 1, 3, 2])
        >>> multiclass_eer(preds, target, num_classes=5, average="macro", thresholds=None)
        tensor(0.4667)
        >>> multiclass_eer(preds, target, num_classes=5, average=None, thresholds=None)
        tensor([0.0000, 0.0000, 0.6667, 0.6667, 1.0000])
        >>> multiclass_eer(preds, target, num_classes=5, average="macro", thresholds=5)
        tensor(0.4667)
        >>> multiclass_eer(preds, target, num_classes=5, average=None, thresholds=5)
        tensor([0.0000, 0.0000, 0.6667, 0.6667, 1.0000])

    )r	   r!   )
r"   r#   r*   r$   r+   r%   r&   r   r   r(   s
             r   Úmulticlass_eerr/   `   s0   € ôD ! ¨°¸ZÈÐR^Ð`mÓn�K€CˆˆaÜ˜˜SÓ!Ð!r   Ú
num_labelsc                 óB   — t        | |||||«      \  }}}t        ||«      S )a 	  Compute Equal Error Rate (EER) for multilabel classification task.

    .. math::
        \text{EER} = \frac{\text{FAR} + (1 - \text{FRR})}{2}, \text{where} \min_t abs(FAR_t-FRR_t)

    The Equal Error Rate (EER) is the point where the False Positive Rate (FPR) and True Positive Rate (TPR) are
    equal, or in practise minimized. A lower EER value signifies higher system accuracy.

    Args:
        preds: Tensor with predictions
        target: Tensor with true labels
        num_labels: Integer specifying the number of labels
        thresholds:
            Can be one of:

            - If set to `None`, will use a non-binned approach where thresholds are dynamically calculated from
              all the data. Most accurate but also most memory consuming approach.
            - If set to an `int` (larger than 1), will use that number of thresholds linearly spaced from
              0 to 1 as bins for the calculation.
            - If set to an `list` of floats, will use the indicated thresholds in the list as bins for the calculation
            - If set to an 1d `tensor` of floats, will use the indicated thresholds in the tensor as
              bins for the calculation.

        ignore_index:
            Specifies a target value that is ignored and does not contribute to the metric calculation
        validate_args: bool indicating if input arguments and tensors should be validated for correctness.
            Set to ``False`` for faster computations.

    Returns:
        A 1d tensor of shape (n_classes, ) will be returned with eer score per label.

    Example:
        >>> from torchmetrics.functional.classification import multilabel_eer
        >>> preds = torch.tensor([[0.75, 0.05, 0.35],
        ...                       [0.45, 0.75, 0.05],
        ...                       [0.05, 0.55, 0.75],
        ...                       [0.05, 0.65, 0.05]])
        >>> target = torch.tensor([[1, 0, 1],
        ...                        [0, 0, 0],
        ...                        [0, 1, 1],
        ...                        [1, 1, 1]])
        >>> multilabel_eer(preds, target, num_labels=3, thresholds=None)
        tensor([0.5000, 0.5000, 0.1667])
        >>> multilabel_eer(preds, target, num_labels=3, thresholds=5)
        tensor([0.5000, 0.7500, 0.1667])

    )r
   r!   )	r"   r#   r0   r$   r%   r&   r   r   r(   s	            r   Úmultilabel_eerr2   ¦   s-   € ôn ! ¨°
¸JÈÐVcÓd�K€CˆˆaÜ˜˜SÓ!Ð!r   Útask)ÚbinaryÚ
multiclassÚ
multilabelc	           	      ó¼  — t        j                  |«      }|t         j                  k(  rt        | ||||«      S |t         j                  k(  r9t        |t        «      st        dt        |«      › d�«      ‚t        | ||||||«      S |t         j                  k(  r8t        |t        «      st        dt        |«      › d�«      ‚t        | |||||«      S t        d|› d�«      ‚)a  Compute Equal Error Rate (EER) metric.

    This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the
    ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of
    :func:`~torchmetrics.functional.classification.binary_eer`,
    :func:`~torchmetrics.functional.classification.multiclass_eer` and
    :func:`~torchmetrics.functional.classification.multilabel_eer` for the specific details of
    each argument influence and examples.

    Args:
        preds: Predictions from model (logits or probabilities)
        target: Ground truth labels
        task: Type of task, either 'binary', 'multiclass' or 'multilabel'
        thresholds: Thresholds used for computing the ROC curve
        num_classes: Number of classes (for multiclass task)
        num_labels: Number of labels (for multilabel task)
        average: Method to average EER over multiple classes/labels
        ignore_index: Specify a target value that is ignored
        validate_args: Bool indicating whether to validate input arguments

    Legacy Example:
        >>> from torchmetrics.functional.classification import eer
        >>> preds = torch.tensor([0.13, 0.26, 0.08, 0.19, 0.34])
        >>> target = torch.tensor([0, 0, 1, 1, 1])
        >>> eer(preds, target, task='binary')
        tensor(0.5833)

        >>> preds = torch.tensor([[0.90, 0.05, 0.05],
        ...                       [0.05, 0.90, 0.05],
        ...                       [0.05, 0.05, 0.90],
        ...                       [0.85, 0.05, 0.10],
        ...                       [0.10, 0.10, 0.80]])
        >>> target = torch.tensor([0, 1, 1, 2, 2])
        >>> eer(preds, target, task='multiclass', num_classes=3, )
        tensor([0.0000, 0.4167, 0.4167])

    z+`num_classes` is expected to be `int` but `z was passed.`z*`num_labels` is expected to be `int` but `zTask z not supported.)r   Úfrom_strÚBINARYr)   Ú
MULTICLASSr   ÚintÚ
ValueErrorÚtyper/   Ú
MULTILABELr2   )	r"   r#   r3   r$   r*   r0   r+   r%   r&   s	            r   Úeerr?   á   sä   € ô` ×&Ñ& tÓ,€DØÔ!×(Ñ(Ò(Ü˜% ¨°\À=ÓQÐQØÔ!×,Ñ,Ò,Ü˜+¤sÔ+ÜÐJÌ4ÐP[ÓK\ÐJ]Ð]jÐkÓlÐlÜ˜e V¨[¸*ÀgÈ|Ð]jÓkÐkØÔ!×,Ñ,Ò,Ü˜*¤cÔ*ÜÐIÌ$ÈzÓJZÐI[Ð[hÐiÓjÐjÜ˜e V¨Z¸À\ÐS`ÓaÐaÜ
�u˜T˜F /Ð2Ó
3Ð3r   )NNT)NNNT)NNNNNT)Útypingr   r   r   r   r   Útyping_extensionsr   Ú*torchmetrics.functional.classification.rocr   r	   r
   Útorchmetrics.utilities.enumsr   r   r!   r;   ÚfloatÚboolr)   r/   r2   r?   © r   r   ú<module>rG      s–  ð÷ )Ñ (ã Ý Ý %÷ñ õ
 <ð+˜Vð +¨&ð +°Vó +ðNØ	ˆv�t˜F‘|Ð#Ñ	$ðNà	ˆv�t˜F‘|Ð#Ñ	$ðNð óNð =AØ"&Øñ0"Øð0"àð0"ð ˜˜s D¨¡K°Ð7Ñ8Ñ9ð0"ð ˜3‘-ð	0"ð
 ð0"ð ó0"ðn =AØ37Ø"&ØñC"ØðC"àðC"ð ðC"ð ˜˜s D¨¡K°Ð7Ñ8Ñ9ð	C"ð
 �gÐ.Ñ/Ñ0ðC"ð ˜3‘-ðC"ð ðC"ð óC"ðT =AØ"&Øñ8"Øð8"àð8"ð ð8"ð ˜˜s D¨¡K°Ð7Ñ8Ñ9ð	8"ð
 ˜3‘-ð8"ð ð8"ð ó8"ð~ =AØ!%Ø $Ø37Ø"&Øñ;4Øð;4àð;4ð Ð6Ñ
7ð;4ð ˜˜s D¨¡K°Ð7Ñ8Ñ9ð	;4ð
 ˜#‘ð;4ð ˜‘ð;4ð �gÐ.Ñ/Ñ0ð;4ð ˜3‘-ð;4ð ð;4ð ˆ6�4˜‘<ÐÑ ô;4r   