Ë
    þÍ:jt0  ã                   ó$  — d dl mZ d dlZd dlmZmZ d dlmZ d dlmZm	Z	m
Z
mZ d dlmZ d dlmZ d dlmZ d	ed
edefd„Zd dedee   ddfd„Zd dededee   ddfd„Zdedededeeef   fd„Z	 	 	 d!dedededee   dedefd„Z	 	 	 d"dededed   dee   ddf
d„Z	 d dedededee   ddf
d„Z	 d#dedededed   deeef   f
d„Z	 	 	 	 d$dededededed   dee   dedefd„Z	 	 	 	 	 d%dededed   dee   deded   dee   dedefd„Zy)&é    )ÚOptionalN)ÚTensorÚtensor)ÚLiteral)Ú_binary_confusion_matrix_formatÚ*_binary_confusion_matrix_tensor_validationÚ#_multiclass_confusion_matrix_formatÚ._multiclass_confusion_matrix_tensor_validation)Únormalize_logits_if_needed)Ú	to_onehot)ÚClassificationTaskNoMultilabelÚmeasureÚtotalÚreturnc                 ó   — | |z  S ©N© )r   r   s     ú�/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/classification/hinge.pyÚ_hinge_loss_computer      s   € Ø�U‰?Ðó    ÚsquaredÚignore_indexc                 ó‚   — t        | t        «      st        d| › �«      ‚|�t        |t        «      st        d|› �«      ‚y y )Nz2Expected argument `squared` to be an bool but got zLExpected argument `ignore_index` to either be `None` or an integer, but got )Ú
isinstanceÚboolÚ
ValueErrorÚint)r   r   s     r   Ú!_binary_hinge_loss_arg_validationr   #   sL   € Ü�gœtÔ$ÜÐMÈgÈYÐWÓXÐXØÐ¬
°<ÄÔ(EÜÐgÐhtÐguÐvÓwÐwð )FÐr   ÚpredsÚtargetc                 ón   — t        | ||«       | j                  «       st        d| j                  › �«      ‚y ©NzdExpected argument `preds` to be floating tensor with probabilities/logits but got tensor with dtype )r   Úis_floating_pointr   Údtype)r   r    r   s      r   Ú$_binary_hinge_loss_tensor_validationr%   *   s?   € Ü.¨u°f¸lÔKØ×"Ñ"Ô$Üð*Ø*/¯+©+¨ð8ó
ð 	
ð %r   c                 ó>  — |j                  «       }t        j                  | «      }| |   ||<   | |     || <   d|z
  }t        j                  |d«      }|r|j	                  d«      }t        |j                  d   |j                  ¬«      }|j                  d¬«      |fS )Né   r   é   ©Údevice©Údim)	r   ÚtorchÚ
zeros_likeÚclampÚpowr   Úshaper*   Úsum)r   r    r   ÚmarginÚmeasuresr   s         r   Ú_binary_hinge_loss_updater5   3   s–   € ð
 �[‰[‹]€FÜ×Ñ˜eÓ$€FØ˜6‘]€Fˆ6�NØ˜f˜W‘~�o€FˆFˆ7�Oà�6‰z€HÜ�{‰{˜8 QÓ'€HáØ—<‘< “?ˆä�6—<‘< ‘?¨6¯=©=Ô9€EØ�<‰<˜Aˆ<Ó Ð%Ð%r   Úvalidate_argsc                 ó–   — |rt        ||«       t        | ||«       t        | |d|d¬«      \  } }t        | ||«      \  }}t	        ||«      S )aó  Compute the mean `Hinge loss`_ typically used for Support Vector Machines (SVMs) for binary tasks.

    .. math::
        \text{Hinge loss} = \max(0, 1 - y \times \hat{y})

    Where :math:`y \in {-1, 1}` is the target, and :math:`\hat{y} \in \mathbb{R}` is the prediction.

    Accepts the following input tensors:

    - ``preds`` (float tensor): ``(N, ...)``. Preds should be a tensor containing probabilities or logits for each
      observation. If preds has values outside [0,1] range we consider the input to be logits and will auto apply
      sigmoid per element.
    - ``target`` (int tensor): ``(N, ...)``. Target should be a tensor containing ground truth labels, and therefore
      only contain {0,1} values (except if `ignore_index` is specified). The value 1 always encodes the positive class.

    Additional dimension ``...`` will be flattened into the batch dimension.

    Args:
        preds: Tensor with predictions
        target: Tensor with true labels
        squared:
            If True, this will compute the squared hinge loss. Otherwise, computes the regular hinge loss.
        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.

    Example:
        >>> from torch import tensor
        >>> from torchmetrics.functional.classification import binary_hinge_loss
        >>> preds = tensor([0.25, 0.25, 0.55, 0.75, 0.75])
        >>> target = tensor([0, 0, 1, 1, 1])
        >>> binary_hinge_loss(preds, target)
        tensor(0.6900)
        >>> binary_hinge_loss(preds, target, squared=True)
        tensor(0.6905)

    g        F)Ú	thresholdr   Úconvert_to_labels)r   r%   r   r5   r   )r   r    r   r   r6   r4   r   s          r   Úbinary_hinge_lossr:   G   sY   € ñZ Ü)¨'°<Ô@Ü,¨U°F¸LÔIÜ3Øˆv °<ÐSXô�M€Eˆ6ô 0°°v¸wÓG�O€HˆeÜ˜x¨Ó/Ð/r   Únum_classesÚmulticlass_mode©úcrammer-singerz
one-vs-allc                 ó’   — t        ||«       t        | t        «      r| dk  rt        d| › �«      ‚d}||vrt        d|› d|› d�«      ‚y )Nr(   zHExpected argument `num_classes` to be an integer larger than 1, but got r=   z1Expected argument `multiclass_mode` to be one of z
, but got ú.)r   r   r   r   )r;   r   r<   r   Ú
allowed_mms        r   Ú%_multiclass_hinge_loss_arg_validationrB   ~   sd   € ô & g¨|Ô<Ü�k¤3Ô'¨;¸ª?ÜÐcÐdoÐcpÐqÓrÐrØ1€JØ˜jÑ(ÜÐLÈZÈLÐXbÐcrÐbsÐstÐuÓvÐvð )r   c                 óp   — t        | |||«       | j                  «       st        d| j                  › �«      ‚y r"   )r
   r#   r   r$   )r   r    r;   r   s       r   Ú(_multiclass_hinge_loss_tensor_validationrD   Œ   sC   € ô 3°5¸&À+È|Ô\Ø×"Ñ"Ô$Üð*Ø*/¯+©+¨ð8ó
ð 	
ð %r   c                 óJ  — t        | d«      } t        |t        d| j                  d   «      «      j	                  «       }|dk(  rD| |   }|t        j                  | |    j                  | j                  d   d«      d¬«      d   z  }n8|j	                  «       }t        j                  | «      }| |   ||<   | |     || <   d|z
  }t        j                  |d«      }|r|j                  d«      }t        |j                  d   |j                  ¬«      }|j                  d¬«      |fS )	NÚsoftmaxr(   r'   r>   r   éÿÿÿÿr+   r)   )r   r   Úmaxr1   r   r-   Úviewr.   r/   r0   r   r*   r2   )r   r    r   r<   r3   r4   r   s          r   Ú_multiclass_hinge_loss_updaterJ   —   s  € ô ' u¨iÓ8€EÜ�vœs 1 e§k¡k°!¡nÓ5Ó6×;Ñ;Ó=€FØÐ*Ò*Ø�v‘ˆØ”%—)‘)˜E 6 '™N×/Ñ/°·±¸A±ÀÓCÈÔKÈAÑNÑN‰à—‘“ˆÜ×!Ñ! %Ó(ˆØ˜v™ˆˆv‰Ø  & ™>˜/ˆ�ˆw‰à�6‰z€HÜ�{‰{˜8 QÓ'€HáØ—<‘< “?ˆä�6—<‘< ‘?¨6¯=©=Ô9€EØ�<‰<˜Aˆ<Ó Ð%Ð%r   c                 óœ   — |rt        ||||«       t        | |||«       t        | ||d¬«      \  } }t        | |||«      \  }}t	        ||«      S )a
  Compute the mean `Hinge loss`_ typically used for Support Vector Machines (SVMs) for multiclass tasks.

    The metric can be computed in two ways. Either, the definition by Crammer and Singer is used:

    .. math::
        \text{Hinge loss} = \max\left(0, 1 - \hat{y}_y + \max_{i \ne y} (\hat{y}_i)\right)

    Where :math:`y \in {0, ..., \mathrm{C}}` is the target class (where :math:`\mathrm{C}` is the number of classes),
    and :math:`\hat{y} \in \mathbb{R}^\mathrm{C}` is the predicted output per class. Alternatively, the metric can
    also be computed in one-vs-all approach, where each class is valued against all other classes in a binary fashion.

    Accepts the following input tensors:

    - ``preds`` (float tensor): ``(N, C, ...)``. Preds should be a tensor containing probabilities or logits for each
      observation. If preds has values outside [0,1] range we consider the input to be logits and will auto apply
      softmax per sample.
    - ``target`` (int tensor): ``(N, ...)``. Target should be a tensor containing ground truth labels, and therefore
      only contain values in the [0, n_classes-1] range (except if `ignore_index` is specified).

    Additional dimension ``...`` will be flattened into the batch dimension.

    Args:
        preds: Tensor with predictions
        target: Tensor with true labels
        num_classes: Integer specifying the number of classes
        squared:
            If True, this will compute the squared hinge loss. Otherwise, computes the regular hinge loss.
        multiclass_mode:
            Determines how to compute the metric
        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.

    Example:
        >>> from torch import tensor
        >>> from torchmetrics.functional.classification import multiclass_hinge_loss
        >>> preds = tensor([[0.25, 0.20, 0.55],
        ...                 [0.55, 0.05, 0.40],
        ...                 [0.10, 0.30, 0.60],
        ...                 [0.90, 0.05, 0.05]])
        >>> target = tensor([0, 1, 2, 0])
        >>> multiclass_hinge_loss(preds, target, num_classes=3)
        tensor(0.9125)
        >>> multiclass_hinge_loss(preds, target, num_classes=3, squared=True)
        tensor(1.1131)
        >>> multiclass_hinge_loss(preds, target, num_classes=3, multiclass_mode='one-vs-all')
        tensor([0.8750, 1.1250, 1.1000])

    F)r9   )rB   rD   r	   rJ   r   )	r   r    r;   r   r<   r   r6   r4   r   s	            r   Úmulticlass_hinge_lossrL   ²   s[   € ñv Ü-¨k¸7ÀOÐUaÔbÜ0°¸ÀÈ\ÔZÜ7¸¸vÀ|ÐglÔm�M€Eˆ6Ü3°E¸6À7ÈOÓ\�O€HˆeÜ˜x¨Ó/Ð/r   Útask)ÚbinaryÚ
multiclassc           	      ó$  — t        j                  |«      }|t         j                  k(  rt        | ||||«      S |t         j                  k(  r9t        |t        «      st        dt        |«      › d�«      ‚t        | ||||||«      S t        d|› �«      ‚)a¿  Compute the mean `Hinge loss`_ typically used for Support Vector Machines (SVMs).

    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'`` or ``'multiclass'``. See the documentation of
    :func:`~torchmetrics.functional.classification.binary_hinge_loss` and
    :func:`~torchmetrics.functional.classification.multiclass_hinge_loss` for the specific details of
    each argument influence and examples.

    Legacy Example:
        >>> from torch import tensor
        >>> target = tensor([0, 1, 1])
        >>> preds = tensor([0.5, 0.7, 0.1])
        >>> hinge_loss(preds, target, task="binary")
        tensor(0.9000)

        >>> target = tensor([0, 1, 2])
        >>> preds = tensor([[-1.0, 0.9, 0.2], [0.5, -1.1, 0.8], [2.2, -0.5, 0.3]])
        >>> hinge_loss(preds, target, task="multiclass", num_classes=3)
        tensor(1.5551)

        >>> target = tensor([0, 1, 2])
        >>> preds = tensor([[-1.0, 0.9, 0.2], [0.5, -1.1, 0.8], [2.2, -0.5, 0.3]])
        >>> hinge_loss(preds, target, task="multiclass", num_classes=3, multiclass_mode="one-vs-all")
        tensor([1.3743, 1.1945, 1.2359])

    z+`num_classes` is expected to be `int` but `z was passed.`zNot handled value: )
r   Úfrom_strÚBINARYr:   Ú
MULTICLASSr   r   r   ÚtyperL   )r   r    rM   r;   r   r<   r   r6   s           r   Ú
hinge_lossrU   õ   s™   € ôH *×2Ñ2°4Ó8€DØÔ-×4Ñ4Ò4Ü  ¨°¸À}ÓUÐUØÔ-×8Ñ8Ò8Ü˜+¤sÔ+ÜÐJÌ4ÐP[ÓK\ÐJ]Ð]jÐkÓlÐlÜ$ U¨F°KÀÈ/Ð[gÐivÓwÐwÜ
Ð*¨4¨&Ð1Ó
2Ð2r   r   )FNF)Fr>   N)r>   )Fr>   NF)NFr>   NT)Útypingr   r-   r   r   Útyping_extensionsr   Ú7torchmetrics.functional.classification.confusion_matrixr   r   r	   r
   Útorchmetrics.utilities.computer   Útorchmetrics.utilities.datar   Útorchmetrics.utilities.enumsr   r   r   r   r   r%   Útupler5   r:   rB   rD   rJ   rL   rU   r   r   r   ú<module>r]      sÐ  ðõ ã ß  Ý %÷ó õ FÝ 1Ý Gð ð °ð ¸6ó ñx¨tð xÀ8ÈCÁ=ð xÐ\`ó xñ
°ð 
Àð 
ÐV^Ð_bÑVcð 
Ðosó 
ð&Øð&àð&ð ð&ð ˆ6�6ˆ>Ñó	&ð. Ø"&Øñ40Øð40àð40ð ð40ð ˜3‘-ð	40ð
 ð40ð ó40ðr Ø?OØ"&ñ	wØðwàðwð Ð;Ñ<ðwð ˜3‘-ð	wð
 
ówð TXñ
Øð
Ø!ð
Ø03ð
ØCKÈCÁ=ð
à	ó
ð @Pñ	&Øð&àð&ð ð&ð Ð;Ñ<ð	&ð
 ˆ6�6ˆ>Ñó&ð> Ø?OØ"&Øñ@0Øð@0àð@0ð ð@0ð ð	@0ð
 Ð;Ñ<ð@0ð ˜3‘-ð@0ð ð@0ð ó@0ðN "&ØØ?OØ"&Øñ+3Øð+3àð+3ð Ð(Ñ
)ð+3ð ˜#‘ð	+3ð
 ð+3ð Ð;Ñ<ð+3ð ˜3‘-ð+3ð ð+3ð ô+3r   