Ë
    þÍ:j¸  ã                   ój   — d dl mZmZmZ d dlZd dlmZ d dlmZmZ d dl	m
Z
 d dlmZ  G d„ de
«      Zy)	é    )ÚAnyÚListÚOptionalN)ÚTensor)Ú_critical_success_index_computeÚ_critical_success_index_update)ÚMetric)Údim_zero_catc            	       óÈ   ‡ — e Zd ZU dZdZeed<   dZeed<   eed<   eed<   eed<   e	e   ed	<   e	e   ed
<   e	e   ed<   dde
dee   deddfˆ fd„Zdededdfd„Zdefd„Zˆ xZS )ÚCriticalSuccessIndexa=  Calculate critical success index (CSI).

    Critical success index (also known as the threat score) is a statistic used weather forecasting that measures
    forecast performance over inputs binarized at a specified threshold. It is defined as:

    .. math:: \text{CSI} = \frac{\text{TP}}{\text{TP}+\text{FN}+\text{FP}}

    Where :math:`\text{TP}`, :math:`\text{FN}` and :math:`\text{FP}` represent the number of true positives, false
    negatives and false positives respectively after binarizing the input tensors.

    Args:
        threshold: Values above or equal to threshold are replaced with 1, below by 0
        keep_sequence_dim: Index of the sequence dimension if the inputs are sequences of images. If specified,
            the score will be calculated separately for each image in the sequence. If ``None``, the score will be
            calculated across all dimensions.

    Example:
        >>> import torch
        >>> from torchmetrics.regression import CriticalSuccessIndex
        >>> x = torch.Tensor([[0.2, 0.7], [0.9, 0.3]])
        >>> y = torch.Tensor([[0.4, 0.2], [0.8, 0.6]])
        >>> csi = CriticalSuccessIndex(0.5)
        >>> csi(x, y)
        tensor(0.3333)

    Example:
        >>> import torch
        >>> from torchmetrics.regression import CriticalSuccessIndex
        >>> x = torch.Tensor([[[0.2, 0.7], [0.9, 0.3]], [[0.2, 0.7], [0.9, 0.3]]])
        >>> y = torch.Tensor([[[0.4, 0.2], [0.8, 0.6]], [[0.4, 0.2], [0.8, 0.6]]])
        >>> csi = CriticalSuccessIndex(0.5, keep_sequence_dim=0)
        >>> csi(x, y)
        tensor([0.3333, 0.3333])

    FÚis_differentiableTÚhigher_is_betterÚhitsÚmissesÚfalse_alarmsÚ	hits_listÚmisses_listÚfalse_alarms_listNÚ	thresholdÚkeep_sequence_dimÚkwargsÚreturnc                 ó  •— t        ‰| �  di |¤Ž t        |«      | _        |r#t	        |t
        «      r|dk  rt        d|› �«      ‚|| _        |€v| j                  dt        j                  d«      d¬«       | j                  dt        j                  d«      d¬«       | j                  dt        j                  d«      d¬«       y | j                  dg d	¬«       | j                  d
g d	¬«       | j                  dg d	¬«       y )Nr   z@Expected keep_sequence_dim to be a non-negative integer but got r   Úsum)ÚdefaultÚdist_reduce_fxr   r   r   Úcatr   r   © )ÚsuperÚ__init__Úfloatr   Ú
isinstanceÚintÚ
ValueErrorr   Ú	add_stateÚtorchÚtensor)Úselfr   r   r   Ú	__class__s       €úp/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/regression/csi.pyr    zCriticalSuccessIndex.__init__G   sÞ   ø€ Ü‰ÑÑ"˜6Ò"Ü˜yÓ)ˆŒá¤jÐ1BÄCÔ&HÐL]Ð`aÒLaÜÐ_Ð`qÐ_rÐsÓtÐtØ!2ˆÔàÐ$Ø�N‰N˜6¬5¯<©<¸«?È5ˆNÔQØ�N‰N˜8¬U¯\©\¸!«_ÈUˆNÔSØ�N‰N˜>´5·<±<À³?ÐSXˆNÕYà�N‰N˜;°À5ˆNÔIØ�N‰N˜=°"ÀUˆNÔKØ�N‰NÐ.¸È5ˆNÕQó    ÚpredsÚtargetc                 óŠ  — t        ||| j                  | j                  «      \  }}}| j                  €@| xj                  |z  c_        | xj                  |z  c_        | xj
                  |z  c_        y| j                  j                  |«       | j                  j                  |«       | j                  j                  |«       y)z*Update state with predictions and targets.N)
r   r   r   r   r   r   r   Úappendr   r   )r(   r,   r-   r   r   r   s         r*   ÚupdatezCriticalSuccessIndex.updateX   s›   € ä%CØ�6˜4Ÿ>™>¨4×+AÑ+Aó&
Ñ"ˆˆf�lð ×!Ñ!Ð)Ø�IŠI˜Ñ�IØ�KŠK˜6Ñ!�KØ×Ò Ñ-Öà�N‰N×!Ñ! $Ô'Ø×Ñ×#Ñ# FÔ+Ø×"Ñ"×)Ñ)¨,Õ7r+   c                 óü   — | j                   €%| j                  }| j                  }| j                  }n?t	        | j
                  «      }t	        | j                  «      }t	        | j                  «      }t        |||«      S )z*Compute critical success index over state.)	r   r   r   r   r
   r   r   r   r   )r(   r   r   r   s       r*   ÚcomputezCriticalSuccessIndex.computef   sh   € à×!Ñ!Ð)Ø—9‘9ˆDØ—[‘[ˆFØ×,Ñ,‰Lä §¡Ó/ˆDÜ! $×"2Ñ"2Ó3ˆFÜ'¨×(>Ñ(>Ó?ˆLÜ.¨t°V¸\ÓJÐJr+   )N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚboolÚ__annotations__r   r   r   r!   r   r#   r   r    r0   r2   Ú__classcell__)r)   s   @r*   r   r      s£   ø… ñ"ðH $Ð�tÓ#Ø!Ð�dÓ!à
ƒLØƒNØÓØ�F‰|ÓØ�f‘ÓØ˜F‘|Ó#ñR %ð R¸HÀS¹Mð RÐ\_ð RÐdhõ Rð"8˜Fð 8¨Fð 8°tó 8ð
K˜÷ 
Kr+   r   )Útypingr   r   r   r&   r   Ú&torchmetrics.functional.regression.csir   r   Útorchmetrics.metricr	   Útorchmetrics.utilitiesr
   r   r   r+   r*   ú<module>r>      s,   ð÷ 'Ñ &ã Ý ç rÝ &Ý /ôXK˜6õ XKr+   