Ë
    þÍ:jq  ã                   ó8  — 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edee	   ddfd	„Z
d
edefd„Zd
ededefd„Zd
edefd„Zdededededef
d„Zdedededeeef   fd„Zdededededeeeef   f
d„Z	 	 ddededed   dee	   deeef   f
d„Zdeddfd„Zy)é    )ÚOptionalN)ÚTensor)ÚLiteral©Úrank_zero_warnÚnan_strategyÚnan_replace_valueÚreturnc                 ó|   — | dvrt        d| › �«      ‚| dk(  r%t        |t        t        f«      st        d|› �«      ‚y y )N©ÚreplaceÚdropzPArgument `nan_strategy` is expected to be one of `['replace', 'drop']`, but got r   zlArgument `nan_replace` is expected to be of a type `int` or `float` when `nan_strategy = 'replace`, but got )Ú
ValueErrorÚ
isinstanceÚfloatÚint)r   r	   s     úz/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/nominal/utils.pyÚ_nominal_input_validationr      sb   € ØÐ.Ñ.ÜØ^Ð_kÐ^lÐmó
ð 	
ð �yÒ ¬Ð4EÌÌsÀ|Ô)TÜðØ(Ð)ð+ó
ð 	
ð *UÐ ó    Úconfmatc                 ó–   — | j                  d«      | j                  d«      }}t        j                  d||«      | j                  «       z  S )zDCompute the expected frequenceis from the provided confusion matrix.é   r   z
r, c -> rc)ÚsumÚtorchÚeinsum)r   Úmargin_sum_rowsÚmargin_sum_colss      r   Ú_compute_expected_freqsr   #   s9   € à'.§{¡{°1£~°w·{±{À1³~�_€OÜ�<‰<˜ o°ÓGÈ'Ï+É+Ë-ÑWÐWr   Úbias_correctionc                 óÄ  — t        | «      }|j                  «       t        |j                  «      z
  |j                  z   dz
  }|dk(  r!t        j                  d| j                  ¬«      S |dk(  rW|rU|| z
  }|j                  «       }| |t        j                  dt        j                  |«      z  |j                  «       «      z  z  } t        j                  | |z
  dz  |z  «      S )z¨Chi-square test of independenc of variables in a confusion matrix table.

    Adapted from: https://github.com/scipy/scipy/blob/v1.9.2/scipy/stats/contingency.py.

    r   r   ç        ©Údeviceg      à?é   )r   Únumelr   ÚshapeÚndimr   Útensorr#   ÚsignÚminimumÚ	ones_likeÚabs)r   r   Úexpected_freqsÚdfÚdiffÚ	directions         r   Ú_compute_chi_squaredr1   )   sÃ   € ô -¨WÓ5€Nà	×	Ñ	Ó	¤# n×&:Ñ&:Ó";Ñ	;¸n×>QÑ>QÑ	QÐTUÑ	U€BØ	ˆQ‚wÜ�|‰|˜C¨¯©Ô7Ð7à	ˆQ‚w‘?Ø Ñ'ˆØ—I‘I“Kˆ	Ø�9œuŸ}™}¨S´5·?±?À9Ó3MÑ-MÈyÏ}É}ËÓ_Ñ_Ñ_ˆä�9‰9�g Ñ.°1Ñ4°~ÑEÓFÐFr   c                 óf   — | | j                  d«      dk7     } | dd…| j                  d«      dk7  f   S )aû  Drop all rows and columns containing only zeros.

    Example:
        >>> from torch import randint
        >>> from torchmetrics.functional.nominal.utils import _drop_empty_rows_and_cols
        >>> matrix = randint(10, size=(4, 3))
        >>> matrix[1, :] = matrix[:, 1] = 0
        >>> matrix
        tensor([[2, 0, 6],
                [0, 0, 0],
                [0, 0, 0],
                [3, 0, 4]])
        >>> _drop_empty_rows_and_cols(matrix)
        tensor([[2, 6],
                [3, 4]])

    r   r   N)r   )r   s    r   Ú_drop_empty_rows_and_colsr3   =   s8   € ð$ �g—k‘k !“n¨Ñ)Ñ*€GØ’1�g—k‘k !“n¨Ñ)Ð)Ñ*Ð*r   Úphi_squaredÚnum_rowsÚnum_colsÚconfmat_sumc                 ó�   — t        j                  t        j                  d| j                  ¬«      | |dz
  |dz
  z  |dz
  z  z
  «      S )z#Compute bias-corrected Phi Squared.r!   r"   r   )r   Úmaxr(   r#   )r4   r5   r6   r7   s       r   Ú_compute_phi_squared_correctedr:   S   sG   € ô �9‰9Ü�‰�S ×!3Ñ!3Ô4Ø˜ 1™¨°A©Ñ6¸;È¹?ÑKÑKóð r   c                 óN   — | | dz
  dz  |dz
  z  z
  }||dz
  dz  |dz
  z  z
  }||fS )z2Compute bias-corrected number of rows and columns.r   r$   © )r5   r6   r7   Úrows_correctedÚcols_correcteds        r   Ú _compute_rows_and_cols_correctedr?   `   sE   € à ¨A¡°!Ñ 3°{ÀQ±Ñ GÑG€NØ ¨A¡°!Ñ 3°{ÀQ±Ñ GÑG€NØ˜>Ð)Ð)r   c                 óH   — t        | |||«      }t        |||«      \  }}|||fS )zBCompute bias-corrected Phi Squared and number of rows and columns.)r:   r?   )r4   r5   r6   r7   Úphi_squared_correctedr=   r>   s          r   Ú_compute_bias_corrected_valuesrB   g   s9   € ô ;¸;ÈÐRZÐ\gÓhÐÜ%EÀhÐPXÐZeÓ%fÑ"€N�NØ  .°.Ð@Ð@r   ÚpredsÚtargetr   c                 óÌ   — |dk(  r"| j                  |«      |j                  |«      fS t        j                  | j                  «       |j                  «       «      }| |    ||    fS )a0  Handle ``NaN`` values in input data.

    If ``nan_strategy = 'replace'``, all ``NaN`` values are replaced with ``nan_replace_value``.
    If ``nan_strategy = 'drop'``, all rows containing ``NaN`` in any of two vectors are dropped.

    Args:
        preds: 1D tensor of categorical (nominal) data
        target: 1D tensor of categorical (nominal) data
        nan_strategy: Indication of whether to replace or drop ``NaN`` values
        nan_replace_value: Value to replace ``NaN`s when ``nan_strategy = 'replace```

    Returns:
        Updated ``preds`` and ``target`` tensors which contain no ``Nan``

    Raises:
        ValueError: If ``nan_strategy`` is not from ``['replace', 'drop']``.
        ValueError: If ``nan_strategy = replace`` and ``nan_replace_value`` is not of a type ``int`` or ``float``.

    r   )Ú
nan_to_numr   Ú
logical_orÚisnan)rC   rD   r   r	   Úrows_contain_nans        r   Ú_handle_nan_in_datarJ   p   sk   € ð2 �yÒ Ø×ÑÐ 1Ó2°F×4EÑ4EÐFWÓ4XÐXÐXÜ×'Ñ'¨¯©«°v·|±|³~ÓFÐØÐ"Ð"Ñ# VÐ-=Ð,=Ñ%>Ð>Ð>r   Úmetric_namec                 ó"   — t        d| › d�«       y )NzUnable to compute zG using bias correction. Please consider to set `bias_correction=False`.r   )rK   s    r   Ú&_unable_to_use_bias_correction_warningrM   �   s   € ÜØ
˜[˜MÐ)pÐqõr   )r   r!   )Útypingr   r   r   Útyping_extensionsr   Útorchmetrics.utilities.printsr   Ústrr   r   r   Úboolr1   r3   r   r:   Útupler?   rB   rJ   rM   r<   r   r   ú<module>rT      s~  ðõ ã Ý Ý %å 8ð	
¨Cð 	
ÀHÈUÁOð 	
ÐX\ó 	
ðX Vð X°ó XðG &ð G¸4ð GÀFó Gð(+ vð +°&ó +ð,
Øð
àð
ð ð
ð ð	
ð
 ó
ð*¨sð *¸cð *ÐPVð *Ð[`ÐagÐioÐaoÑ[pó *ðAØðAØ#&ðAØ25ðAØDJðAà
ˆ6�6˜6Ð!Ñ"óAð 09Ø),ñ	?Øð?àð?ð Ð+Ñ,ð?ð   ‘ð	?ð
 ˆ6�6ˆ>Ñó?ð>¸ð Àô r   