Ë
    þÍ:jÈ=  ã                   ó  — 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	 d dl
mZ d dlmZmZmZ d dlmZ  G d	„ d
e«      Z G d„ de«      Zdededeeef   fd„Zdedededefd„Zdededefd„Zdedededefd„Zdededefd„Zd5dededefd„Zdedeeeef   fd„Zdedededeeeee   ee   ee   ee   ee   ee   ef	   fd„Zdededededed ed!ee   d"ee   dedefd#„Zd$edefd%„Z ded ed!ee   d&ee   d'ee   d"ee   d(ee   d)ee   ded*ee   defd+„Z!	 	 	 d6deded,eee      d-eee      d.edeee   ee   f   fd/„Z"	 d7dededed*ee   deeee   f   f
d0„Z#	 	 	 d8dededed1   d2ed*eed3      deeeeef   f   fd4„Z$y)9é    )ÚListÚOptionalÚUnionN)ÚTensor)ÚLiteral)Ú _check_data_shape_to_num_outputs)Ú_check_same_shape)Ú	_bincountÚ_cumsumÚdim_zero_cat)ÚEnumStrc                   ó2   — e Zd ZdZdZdZdZedefd„«       Z	y)Ú_MetricVariantzEnumerate for metric variants.ÚaÚbÚcÚreturnc                   ó   — y)NÚvariant© r   ó    ú/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/regression/kendall.pyÚ_namez_MetricVariant._name!   s   € àr   N)
Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚAÚBÚCÚstaticmethodÚstrr   r   r   r   r   r      s-   „ Ù(à€AØ€AØ€Aàð�3ò ó ñr   r   c                   ó2   — e Zd ZdZdZdZdZedefd„«       Z	y)Ú_TestAlternativez'Enumerate for test alternative options.ú	two-sidedÚlessÚgreaterr   c                   ó   — y)NÚalternativer   r   r   r   r   z_TestAlternative._name-   s   € àr   N)
r   r   r   r   Ú	TWO_SIDEDÚLESSÚGREATERr!   r"   r   r   r   r   r$   r$   &   s-   „ Ù1à€IØ€DØ€Gàð�3ò ó ñr   r$   ÚxÚyr   c                 ó  — t        j                  |«      }| j                  |j                  }} | j                  «       \  } }t	        | j
                  d   «      D ]  }||   ||      ||<   Œ | j                  |j                  fS )zBSort sequences in an ascent order according to the sequence ``x``.r   )ÚtorchÚcloneÚTÚsortÚrangeÚshape)r-   r.   ÚpermÚis       r   Ú_sort_on_first_sequencer8   2   ss   € ô 	�‰�A‹€AØ�3‰3�—‘€q€AØ�f‰f‹h�G€A€tÜ�1—7‘7˜1‘:Óò ˆØ�‰t�D˜‘G‰}ˆˆ!Šðà�3‰3�—‘ˆ8€Or   r7   c                 óš   — t        j                  | |   | |dz   d k  ||   ||dz   d k  «      j                  d«      j                  d«      S )z>Count a total number of concordant pairs in a single sequence.é   Nr   )r0   Úlogical_andÚsumÚ	unsqueeze©r-   r.   r7   s      r   Ú_concordant_element_sumr?   =   sR   € ä×Ñ˜Q˜q™T A q¨1¡u i LÑ0°!°A±$¸¸AÀ¹E¸9¸Ñ2EÓF×JÑJÈ1ÓM×WÑWÐXYÓZÐZr   ÚpredsÚtargetc           
      ó²   — t        j                  t        | j                  d   «      D �cg c]  }t	        | ||«      ‘Œ c}«      j                  d«      S c c}w )z<Count a total number of concordant pairs in given sequences.r   )r0   Úcatr4   r5   r?   r<   ©r@   rA   r7   s      r   Ú_count_concordant_pairsrE   B   óC   € ä�9‰9ÌÈuÏ{É{Ð[\É~ÓI^Ö_ÀAÔ-¨e°V¸QÕ?Ò_Ó`×dÑdÐefÓgÐgùÒ_ó   ªAc           
      ó  — t        j                  t        j                  | |   | |dz   d kD  ||   ||dz   d k  «      t        j                  | |   | |dz   d k  ||   ||dz   d kD  «      «      j                  d«      j	                  d«      S )z?Count a total number of discordant pairs in a single sequences.r:   Nr   )r0   Ú
logical_orr;   r<   r=   r>   s      r   Ú_discordant_element_sumrJ   G   s—   € ô 	×ÑÜ×Ñ˜a ™d Q¨¨A© y \Ñ1°1°Q±4¸!¸QÀ¹U¸I¸,Ñ3FÓGÜ×Ñ˜a ™d Q¨¨A© y \Ñ1°1°Q±4¸!¸QÀ¹U¸I¸,Ñ3FÓGó	
÷ 
‰ˆQ‹ß	‰�1‹ðr   c           
      ó²   — t        j                  t        | j                  d   «      D �cg c]  }t	        | ||«      ‘Œ c}«      j                  d«      S c c}w )z<Count a total number of discordant pairs in given sequences.r   )r0   rC   r4   r5   rJ   r<   rD   s      r   Ú_count_discordant_pairsrL   S   rF   rG   r3   c           	      ó0  — |r| j                  d¬«      j                  } t        j                  d| j                  d   t        j
                  | j                  ¬«      }t        t        j                  || dd | dd k7  j                  «       gd¬«      d¬«      S )z&Convert a sequence to the rank tensor.r   )Údimr:   ©ÚdtypeÚdeviceNéÿÿÿÿ)
r3   Úvaluesr0   Úzerosr5   Úint32rQ   r   rC   Úint)r-   r3   Ú_oness      r   Ú_convert_sequence_to_dense_rankrX   X   sv   € ñ Ø�F‰F�qˆF‹M× Ñ ˆÜ�K‰K˜˜1Ÿ7™7 1™:¬U¯[©[ÀÇÁÔJ€EÜ”5—9‘9˜e a¨¨ e¨q°°"¨v¡o×%:Ñ%:Ó%<Ð=À1ÔEÈ1ÔMÐMr   c                 óŠ  — t        j                  | j                  d   | j                  | j                  ¬«      }t        j                  | j                  d   | j                  | j                  ¬«      }t        j                  | j                  d   | j                  | j                  ¬«      }t        | j                  d   «      D ]y  }t        | dd…|f   «      }||dkD     }||dz
  z  dz  j                  «       ||<   ||dz
  z  |dz
  z  j                  «       ||<   ||dz
  z  d|z  dz   z  j                  «       ||<   Œ{ |||fS )zWGet a total number of ties and staistics for p-value calculation for  a given sequence.r:   rO   Né   ç      ð?é   )r0   rT   r5   rP   rQ   r4   r
   r<   )r-   ÚtiesÚties_p1Úties_p2rN   Ún_tiess         r   Ú	_get_tiesra   a   s!  € ä�;‰;�q—w‘w˜q‘z¨¯©¸¿¹ÔB€DÜ�k‰k˜!Ÿ'™' !™*¨A¯G©G¸A¿H¹HÔE€GÜ�k‰k˜!Ÿ'™' !™*¨A¯G©G¸A¿H¹HÔE€GÜ�Q—W‘W˜Q‘ZÓ ò JˆÜ˜1šQ ˜V™9Ó%ˆØ˜ ™
Ñ#ˆØ˜v¨™zÑ*¨aÑ/×4Ñ4Ó6ˆˆS‰	Ø &¨3¡,Ñ/°6¸A±:Ñ>×CÑCÓEˆ�‰Ø &¨3¡,Ñ/°1°v±:À±>ÑB×GÑGÓIˆ�ŠðJð �˜'Ð!Ð!r   r   c           	      ól  — t        | |«      \  } }t        | |«      }t        | |«      }t        j                  | j
                  d   | j                  ¬«      }dx}}dx}x}	x}
}|t        j                  k7  r6t        | «      } t        |d¬«      }t        | «      \  }}}	t        |«      \  }}
}|||||	||
||f	S )z,Obtain statistics to calculate metric value.r   )rQ   NT)r3   )r8   rE   rL   r0   Útensorr5   rQ   r   r   rX   ra   )r@   rA   r   Úconcordant_pairsÚdiscordant_pairsÚn_totalÚ
preds_tiesÚtarget_tiesÚpreds_ties_p1Úpreds_ties_p2Útarget_ties_p1Útarget_ties_p2s               r   Ú_get_metric_metadatarm   p   sÓ   € ô ,¨E°6Ó:�M€Eˆ6ä.¨u°fÓ=ÐÜ.¨u°fÓ=Ðä�l‰l˜5Ÿ;™; q™>°%·,±,Ô?€GØ#Ð#€J�ØFJÐJ€MÐJ�MÐJ N°^Ø”.×"Ñ"Ò"Ü/°Ó6ˆÜ0°¸dÔCˆÜ3<¸UÓ3CÑ0ˆ
�M =Ü6?ÀÓ6GÑ3ˆ�^ ^àØØØØØØØØð
ð 
r   rd   re   Úcon_min_dis_pairsrf   rg   rh   c	           	      óF  — |t         j                  k(  r|||z   z  S |t         j                  k(  rŠ||dz
  z  dz  }	|€,t        j                  d|	j
                  |	j                  ¬«      }|€,t        j                  d|	j
                  |	j                  ¬«      }|	|z
  |	|z
  z  }
|t        j                  |
«      z  S t        j                  | j                  D �cg c]  }t        |j                  «       «      ‘Œ c}| j
                  | j                  ¬«      }t        j                  |j                  D �cg c]  }t        |j                  «       «      ‘Œ c}|j
                  |j                  ¬«      }t        j                  ||«      }d|z  |dz
  |z  |dz  z  z  S c c}w c c}w )z-Calculate Kendall's tau from metric metadata.r:   rZ   ç        rO   )r   r   r   r0   rc   rP   rQ   Úsqrtr2   ÚlenÚuniqueÚminimum)r@   rA   rd   re   rn   rf   rg   rh   r   Útotal_combinationsÚdenominatorÚpÚpreds_uniqueÚtÚtarget_uniqueÚmin_classess                   r   Ú_calculate_taur|   ˜   sd  € ð ”.×"Ñ"Ò"Ø Ð$4Ð7GÑ$GÑHÐHØ”.×"Ñ"Ò"Ø%,°¸!±Ñ%<ÀÑ%AÐØÐÜŸ™ cÐ1C×1IÑ1IÐRd×RkÑRkÔlˆJØÐÜŸ,™, sÐ2D×2JÑ2JÐSe×SlÑSlÔmˆKØ)¨JÑ6Ð;MÐP[Ñ;[Ñ\ˆØ ¤5§:¡:¨kÓ#:Ñ:Ð:ä—<‘<¸%¿'¹'Ö B°Q¤ Q§X¡X£Z¥Ò BÈ%Ï+É+Ð^c×^jÑ^jÔk€LÜ—L‘L¸6¿8¹8Ö!D°a¤# a§h¡h£j¥/Ò!DÈFÏLÉLÐag×anÑanÔo€MÜ—-‘- ¨mÓ<€KØÐ Ñ  [°1¡_¸Ñ$CÀgÈqÁjÑ$PÑQÐQùò !CùÚ!Ds   Ã FÄ1 FÚt_valuec                 óÔ  — | }t         j                  j                  j                  t        j                  dg«      j                  |«      t        j                  dg«      j                  |«      «      }| j                  «       }| j                  «       } |j                  | «      }|j                  | t        j                  t        d«      |j                  |j                  ¬«      «      S )z¾Obtain p-value for a given Tensor of t-values. Handle ``nan`` which cannot be passed into torch distributions.

    When t-value is ``nan``, a resulted p-value should be alson ``nan``.

    rp   r[   ÚnanrO   )r0   ÚdistributionsÚnormalÚNormalrc   ÚtoÚisnanÚ
nan_to_numÚcdfÚwhereÚfloatrP   rQ   )r}   rQ   Únormal_distÚis_nanÚp_values        r   Ú"_get_p_value_for_t_value_from_distrŒ   µ   s¯   € ð €FÜ×%Ñ%×,Ñ,×3Ñ3´E·L±LÀ#ÀÓ4G×4JÑ4JÈ6Ó4RÔTY×T`ÑT`ÐbeÐafÓTg×TjÑTjÐkqÓTrÓs€Kà�]‰]‹_€FØ× Ñ Ó"€GØ�o‰o˜gÓ&€GØ�=‰=˜&˜¤%§,¡,¬u°U«|À7Ç=Á=ÐY`×YgÑYgÔ"hÓiÐir   ri   rj   rk   rl   r)   c
                 ó:  — ||dz
  z  d|z  dz   z  }
|t         j                  k(  rd| z  t        j                  |
dz  «      z  }ne||dz
  z  }|
|�|ndz
  |�|ndz
  dz  }|d|�|ndz  |�|ndz  |z  z  }||�|nd|�|ndz  d|z  |dz
  z  z  z  }| t        j                  |«      z  }|	t        j
                  k(  rt        j                  |«      }|	t        j
                  t        j                  fv r|dz  }t        |«      }|	t        j
                  k(  r|dz  }|S )	z9Calculate p-value for Kendall's tau from metric metadata.r:   rZ   r\   é   r   é   é	   rR   )	r   r   r0   rq   r$   r*   Úabsr,   rŒ   )rn   rf   rg   ri   rj   rh   rk   rl   r   r)   Út_value_denominator_baser}   ÚmÚt_value_denominatorr‹   s                  r   Ú_calculate_p_valuer•   Ä   sx  € ð  '¨'°A©+Ñ6¸!¸g¹+È¹/ÑJÐØ”.×"Ñ"Ò"ØÐ'Ñ'¬%¯*©*Ð5MÐPQÑ5QÓ*RÑR‰à�w ‘{Ñ#ˆà$Ø -Ð 9‰}¸qñBà!/Ð!;‰~ÀñDð ñ	'Ðð
 	Ø˜zÐ5‘¸1Ñ=ÐP[ÐPgÁÐmnÑoØñ ñ 	Ðð 	Ø+Ð7‰]¸QØ!/Ð!;‰~ÀñDà�1‰u˜ !™Ñ$ñ&ñ	
Ðð
 $¤e§j¡jÐ1DÓ&EÑEˆàÔ&×0Ñ0Ò0Ü—)‘)˜GÓ$ˆØÔ'×1Ñ1Ô3C×3KÑ3KÐLÑLØ�2‰ˆÜ0°Ó9€GØÔ&×0Ñ0Ò0Ø�1‰ˆØ€Nr   Úconcat_predsÚconcat_targetÚnum_outputsc                 óæ   — |xs g }|xs g }t        | |«       t        | ||«       |dk(  r"| j                  d«      } |j                  d«      }|j                  | «       |j                  |«       ||fS )aÌ  Update variables required to compute Kendall rank correlation coefficient.

    Args:
        preds: Sequence of data
        target: Sequence of data
        concat_preds: List of batches of preds sequence to be concatenated
        concat_target: List of batches of target sequence to be concatenated
        num_outputs: Number of outputs in multioutput setting

    Raises:
        RuntimeError: If ``preds`` and ``target`` do not have the same shape

    r:   )r	   r   r=   Úappend)r@   rA   r–   r—   r˜   s        r   Ú_kendall_corrcoef_updater›   ï   s{   € ð(  Ò% 2€LØ!Ò' R€Mä�e˜VÔ$Ü$ U¨F°KÔ@à�aÒØ—‘ Ó"ˆØ×!Ñ! !Ó$ˆà×Ñ˜ÔØ×Ñ˜Ô à˜Ð&Ð&r   c                 ó$  — t        | ||«      \	  }}}}}}	}
}}||z
  }t        | |||||||	|«	      }|rt        ||||||	|
|||«
      nd}|j                  d   dk(  r$|j	                  «       }|�|j	                  «       nd}|j                  dd«      |fS )a/  Compute Kendall rank correlation coefficient, and optionally p-value of corresponding statistical test.

    Args:
        Args:
        preds: Sequence of data
        target: Sequence of data
        variant: Indication of which variant of Kendall's tau to be used
        alternative: Alternative hypothesis for for t-test. Possible values:
            - 'two-sided': the rank correlation is nonzero
            - 'less': the rank correlation is negative (less than zero)
            - 'greater':  the rank correlation is positive (greater than zero)

    Nr   r:   rR   )rm   r|   r•   r5   ÚsqueezeÚclamp)r@   rA   r   r)   rd   re   rg   ri   rj   rh   rk   rl   rf   rn   Útaur‹   s                   r   Ú_kendall_corrcoef_computer      sÜ   € ô: 	˜U F¨GÓ4ñ
ØØØØØØØØØà(Ð+;Ñ;Ðä
ØˆvÐ'Ð)9Ð;LÈgÐWaÐcnÐpwó€Cñ  ô 	ØØØØØØØØØØô	
ð ð ð$ ‡y�y��|�qÒØ�k‰k‹mˆØ'.Ð':�'—/‘/Ô#Àˆà�9‰9�R˜Ó˜WÐ$Ð$r   )r   r   r   Út_test)r%   r&   r'   c                 ó¨  — t        |t        «      st        dt        |«      › d�«      ‚|r|€t        d«      ‚t        j                  t        |«      «      }|rt        j                  t        |«      «      nd}t        | |g g | j                  dk(  rdn| j                  d   ¬«      \  }}t        t        |«      t        |«      ||«      \  }	}
|
�|	|
fS |	S )a   Compute `Kendall Rank Correlation Coefficient`_.

    .. math::
        tau_a = \frac{C - D}{C + D}

    where :math:`C` represents concordant pairs, :math:`D` stands for discordant pairs.

    .. math::
        tau_b = \frac{C - D}{\sqrt{(C + D + T_{preds}) * (C + D + T_{target})}}

    where :math:`C` represents concordant pairs, :math:`D` stands for discordant pairs and :math:`T` represents
    a total number of ties.

    .. math::
        tau_c = 2 * \frac{C - D}{n^2 * \frac{m - 1}{m}}

    where :math:`C` represents concordant pairs, :math:`D` stands for discordant pairs, :math:`n` is a total number
    of observations and :math:`m` is a ``min`` of unique values in ``preds`` and ``target`` sequence.

    Definitions according to Definition according to `The Treatment of Ties in Ranking Problems`_.

    Args:
        preds: Sequence of data of either shape ``(N,)`` or ``(N,d)``
        target: Sequence of data of either shape ``(N,)`` or ``(N,d)``
        variant: Indication of which variant of Kendall's tau to be used
        t_test: Indication whether to run t-test
        alternative: Alternative hypothesis for t-test. Possible values:
            - 'two-sided': the rank correlation is nonzero
            - 'less': the rank correlation is negative (less than zero)
            - 'greater':  the rank correlation is positive (greater than zero)

    Return:
        Correlation tau statistic
        (Optional) p-value of corresponding statistical test (asymptotic)

    Raises:
        ValueError: If ``t_test`` is not of a type bool
        ValueError: If ``t_test=True`` and ``alternative=None``

    Example (single output regression):
        >>> from torchmetrics.functional.regression import kendall_rank_corrcoef
        >>> preds = torch.tensor([2.5, 0.0, 2, 8])
        >>> target = torch.tensor([3, -0.5, 2, 1])
        >>> kendall_rank_corrcoef(preds, target)
        tensor(0.3333)

    Example (multi output regression):
        >>> from torchmetrics.functional.regression import kendall_rank_corrcoef
        >>> preds = torch.tensor([[2.5, 0.0], [2, 8]])
        >>> target = torch.tensor([[3, -0.5], [2, 1]])
        >>> kendall_rank_corrcoef(preds, target)
        tensor([1., 1.])

    Example (single output regression with t-test)
        >>> from torchmetrics.functional.regression import kendall_rank_corrcoef
        >>> preds = torch.tensor([2.5, 0.0, 2, 8])
        >>> target = torch.tensor([3, -0.5, 2, 1])
        >>> kendall_rank_corrcoef(preds, target, t_test=True, alternative='two-sided')
        (tensor(0.3333), tensor(0.4969))

    Example (multi output regression with t-test):
        >>> from torchmetrics.functional.regression import kendall_rank_corrcoef
        >>> preds = torch.tensor([[2.5, 0.0], [2, 8]])
        >>> target = torch.tensor([[3, -0.5], [2, 1]])
        >>> kendall_rank_corrcoef(preds, target, t_test=True, alternative='two-sided')
            (tensor([1., 1.]), tensor([nan, nan]))

    z>Argument `t_test` is expected to be of a type `bool`, but got ú.NzCArgument `alternative` is required if `t_test=True` but got `None`.r:   rR   )r˜   )Ú
isinstanceÚboolÚ
ValueErrorÚtyper   Úfrom_strr"   r$   r›   Úndimr5   r    r   )r@   rA   r   r¡   r)   Ú_variantÚ_alternativeÚ_predsÚ_targetrŸ   r‹   s              r   Úkendall_rank_corrcoefr®   O  sÚ   € ôV �fœdÔ#ÜÐYÔZ^Ð_eÓZfÐYgÐghÐiÓjÐjÙ�+Ð%ÜÐ^Ó_Ð_ä×&Ñ&¤s¨7£|Ó4€HÙBHÔ#×,Ñ,¬S°Ó-=Ô>Èd€Lä.Øˆv�r˜2°·
±
¸a²©1ÀUÇ[Á[ÐQSÁ_ô�O€FˆGô -Ü�VÓÜ�WÓØØó	�L€Cˆð ÐØ�Gˆ|ÐØ€Jr   )F)NNr:   )N)r   Fr%   )%Útypingr   r   r   r0   r   Útyping_extensionsr   Ú(torchmetrics.functional.regression.utilsr   Útorchmetrics.utilities.checksr	   Útorchmetrics.utilities.datar
   r   r   Útorchmetrics.utilities.enumsr   r   r$   Útupler8   rV   r?   rE   rJ   rL   r¥   rX   ra   rm   r|   rŒ   r•   r›   r    r®   r   r   r   ú<module>r¶      sÚ  ð÷ )Ñ (ã Ý Ý %å UÝ ;ß HÑ HÝ 0ô	�Wô 	ô	�wô 	ð˜vð ¨&ð °U¸6À6¸>Ñ5Jó ð[˜vð [¨&ð [°Sð [¸Vó [ð
h 6ð h°6ð h¸fó hð
	˜vð 	¨&ð 	°Sð 	¸Vó 	ðh 6ð h°6ð h¸fó hñ
N vð N°Tð NÀfó Nð"�ð "˜E &¨&°&Ð"8Ñ9ó "ð%Øð%Ø!ð%Ø,:ð%à
Ø
Ø
ØˆVÑØˆVÑØˆVÑØˆVÑØˆVÑØˆVÑØ
ðñ
ó%ðPRØðRàðRð ðRð ð	Rð
 ðRð ðRð ˜Ñ ðRð ˜&Ñ!ðRð ðRð óRð:j°ð j¸6ó jð(Øð(àð(ð ˜Ñ ð(ð ˜FÑ#ð	(ð
 ˜FÑ#ð(ð ˜&Ñ!ð(ð ˜VÑ$ð(ð ˜VÑ$ð(ð ð(ð Ð*Ñ+ð(ð ó(ð\ ,0Ø,0Øñ!'Øð!'àð!'ð ˜4 ™<Ñ(ð!'ð ˜D ™LÑ)ð	!'ð
 ð!'ð ˆ4�‰<˜˜f™Ð%Ñ&ó!'ðP /3ñ	9%Øð9%àð9%ð ð9%ð Ð*Ñ+ð	9%ð
 ˆ6�8˜FÑ#Ð#Ñ$ó9%ð~ '*ØØEPñ_Øð_àð_ð �]Ñ#ð_ð ð	_ð
 ˜'Ð"@ÑAÑBð_ð ˆ6�5˜ ˜Ñ(Ð(Ñ)ô_r   