Ë
    þÍ:j÷  ã                   ó¬   — d dl mZ d dlZd dlmZ d dlmZ d dlmZ 	 ddededed	ee	   d
e
eeef   f
d„Zdededed
efd„Z	 ddededed	ee	   d
ef
d„Zy)é    )ÚOptionalN)ÚTensor)Ú_check_same_shape©Ú_safe_divideÚpredsÚtargetÚ	thresholdÚkeep_sequence_dimÚreturnc                 ó   ‡— t        | |«       ‰€d}nYd‰cxk  r| j                  k  sn t        d| j                  › d‰› �«      ‚t        ˆfd„t	        | j                  «      D «       «      }| |k\  j                  «       }||k\  j                  «       }‰€yt        j                  ||z  «      j                  «       }t        j                  ||z  |z  «      j                  «       }t        j                  ||z  |z  «      j                  «       }	n~t        j                  ||z  |¬«      j                  «       }t        j                  ||z  |z  |¬«      j                  «       }t        j                  ||z  |z  |¬«      j                  «       }	|||	fS )a%  Update and return variables required to compute Critical Success Index. Checks for same shape of tensors.

    Args:
        preds: Predicted tensor
        target: Ground truth tensor
        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.

    Nr   z.Expected keep_sequence dim to be in range [0, z
] but got c              3   ó.   •K  — | ]  }|‰k7  sŒ	|–— Œ y ­w©N© )Ú.0Úir   s     €ú{/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/regression/csi.pyú	<genexpr>z1_critical_success_index_update.<locals>.<genexpr>,   s   øè ø€ ÒP˜q¸Ð>OÓ9OœÑPùs   ƒ
Ž)Údim)	r   ÚndimÚ
ValueErrorÚtupleÚrangeÚboolÚtorchÚsumÚint)
r   r	   r
   r   Úsum_dimsÚ	preds_binÚ
target_binÚhitsÚmissesÚfalse_alarmss
      `      r   Ú_critical_success_index_updater$      sb  ø€ ô �e˜VÔ$àÐ Ø‰ØÐ#Ô0 e§j¡jÔ0ÜÐIÈ%Ï*É*ÈÐU_Ð`qÐ_rÐsÓtÐtäÓP¤E¨%¯*©*Ó$5ÔPÓPˆð ˜)Ñ#×)Ñ)Ó+€IØ˜IÑ%×+Ñ+Ó-€JàÐ Ü�y‰y˜ ZÑ/Ó0×4Ñ4Ó6ˆÜ—‘˜I¨
Ñ2°jÑ@ÓA×EÑEÓGˆÜ—y‘y )¨jÑ"8¸IÑ!EÓF×JÑJÓL‰ä�y‰y˜ ZÑ/°XÔ>×BÑBÓDˆÜ—‘˜I¨
Ñ2°jÑ@ÀhÔO×SÑSÓUˆÜ—y‘y )¨jÑ"8¸IÑ!EÈ8ÔT×XÑXÓZˆØ�˜Ð%Ð%ó    r!   r"   r#   c                 ó&   — t        | | |z   |z   «      S )aÚ  Compute critical success index.

    Args:
        hits: Number of true positives after binarization
        misses: Number of false negatives after binarization
        false_alarms: Number of false positives after binarization

    Returns:
        If input tensors are 5-dimensional and ``keep_sequence_dim=True``, the metric returns a ``(S,)`` vector
        with CSI scores for each image in the sequence. Otherwise, it returns a scalar tensor with the CSI score.

    r   )r!   r"   r#   s      r   Ú_critical_success_index_computer'   =   s   € ô ˜˜d V™m¨lÑ:Ó;Ð;r%   c                 ó@   — t        | |||«      \  }}}t        |||«      S )aY  Compute critical success index.

    Args:
        preds: Predicted tensor
        target: Ground truth tensor
        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.

    Returns:
        If ``keep_sequence_dim`` is specified, the metric returns a vector of  with CSI scores for each image
        in the sequence. Otherwise, it returns a scalar tensor with the CSI score.

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

    Example:
        >>> import torch
        >>> from torchmetrics.functional.regression import critical_success_index
        >>> 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]]])
        >>> critical_success_index(x, y, 0.5, keep_sequence_dim=0)
        tensor([0.3333, 0.3333])

    )r$   r'   )r   r	   r
   r   r!   r"   r#   s          r   Úcritical_success_indexr)   M   s-   € ôD "@ÀÀvÈyÐZkÓ!lÑ€Dˆ&�,Ü*¨4°¸ÓFÐFr%   r   )Útypingr   r   r   Útorchmetrics.utilities.checksr   Útorchmetrics.utilities.computer   Úfloatr   r   r$   r'   r)   r   r%   r   ú<module>r.      s¾   ðõ ã Ý å ;Ý 7ð Y]ñ#&Øð#&Ø!ð#&Ø.3ð#&ØHPÐQTÉð#&à
ˆ6�6˜6Ð!Ñ"ó#&ðL<¨&ð <¸&ð <ÐPVð <Ð[aó <ð" Y]ñ#GØð#GØ!ð#GØ.3ð#GØHPÐQTÉð#Gàô#Gr%   