Ë
    ýÿæi   ã                   ón  — d dl mZ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defd	„Z	dededefd
„Z
	 d%dededeeed   f   defd„Z	 d&dedee   dedededededefd„Zdededeeef   fd„Zd'dededededef
d„Zd(dedededefd„Zd(dedededefd„Zdedededefd „Zd!ed"eed#      defd$„Zy))é    )ÚOptionalÚUnionN)ÚTensor)ÚLiteral)Úrank_zero_warnÚxÚyÚreturnc                 ó  — | j                   t        j                  k(  s|j                   t        j                  k(  r9| j                  «       |j                  j                  «       z  j                  «       S | |j                  z  S )zSafe calculation of matrix multiplication.

    If input is float16, will cast to float32 for computation and back again.

    )ÚdtypeÚtorchÚfloat16ÚfloatÚTÚhalf©r   r	   s     ús/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchmetrics/utilities/compute.pyÚ_safe_matmulr      sT   € ð 	‡w�w”%—-‘-Ò 1§7¡7¬e¯m©mÒ#;Ø—‘“	˜AŸC™CŸI™I›KÑ'×-Ñ-Ó/Ð/Øˆq�s‰s‰7€Nó    c                 óF   — | t        j                  |«      z  }d|| dk(  <   |S )z¦Compute x * log(y). Returns 0 if x=0.

    Example:
        >>> import torch
        >>> x = torch.zeros(1)
        >>> _safe_xlogy(x, 1/x)
        tensor([0.])

    ç        r   )r   Úlog)r   r	   Úress      r   Ú_safe_xlogyr   "   s(   € ð Œe�i‰i˜‹lÑ
€CØ€CˆˆQ‰�KØ€Jr   ÚnumÚdenomÚzero_division)ÚwarnÚnanc                 óâ  — | j                  «       r| n| j                  «       } |j                  «       r|n|j                  «       }t        |t        t        f«      s|dk(  r{|dk(  r#t	        j
                  |dk(  «      rt        d«       |dk(  rdn|}t	        j                  d|| j                  | j                  ¬«      }t	        j                  |dk7  | |z  |«      S t	        j                  | |«      S )a  Safe division, by preventing division by zero.

    Function will cast to float if input is not already to secure backwards compatibility.

    Args:
        num: numerator tensor
        denom: denominator tensor, which may contain zeros
        zero_division: value to replace elements divided by zero

    Example:
        >>> import torch
        >>> num = torch.tensor([1.0, 2.0, 3.0])
        >>> denom = torch.tensor([0.0, 1.0, 2.0])
        >>> _safe_divide(num, denom)
        tensor([0.0000, 2.0000, 1.5000])

    r   r   z:Detected zero division in _safe_divide. Setting 0/0 to 0.0r   © )r   Údevice)Úis_floating_pointr   Ú
isinstanceÚintr   Úanyr   Úfullr   r"   ÚwhereÚtrue_divide)r   r   r   Úzero_division_tensors       r   Ú_safe_divider+   1   sÆ   € ð, ×&Ñ&Ô(‰#¨c¯i©i«k€CØ×,Ñ,Ô.‰E°E·K±K³M€EÜ�-¤%¬ Ô.°-À6Ò2IØ˜FÒ"¤u§y¡y°¸!±Ô'<ÜÐWÔXØ,°Ò6™¸MˆÜ$Ÿz™z¨"¨mÀ3Ç9Á9ÐUX×U_ÑU_Ô`ÐÜ�{‰{˜5 A™: s¨U¡{Ð4HÓIÐIÜ×Ñ˜S %Ó(Ð(r   ÚscoreÚaverageÚ
multilabelÚtpÚfpÚfnÚtop_kc                 óì   — |�|dk(  r| S |dk(  r||z   }n2t        j                  | «      }|sd||dk(  r||z   |z   dk(  n||z   dk(  <   t        || z  |j                  dd¬«      «      j                  d«      S )	NÚnoneÚweightedr   é   r   éÿÿÿÿT)Úkeepdim)r   Ú	ones_liker+   Úsum)r,   r-   r.   r/   r0   r1   r2   Úweightss           r   Ú_adjust_weights_safe_divider<   R   s‡   € ð €˜' VÒ+ØˆØ�*ÒØ�r‘'‰ä—/‘/ %Ó(ˆÙØILˆG¨°!ª�B˜‘G˜b‘L AÒ%¸¸b¹ÀA¹ÑFÜ˜ %™¨¯©°RÀ¨Ó)FÓG×KÑKÈBÓOÐOr   c                 ó°  — | j                   dkD  r| j                  «       n| } |j                   dkD  r|j                  «       n|}| j                   dkD  s|j                   dkD  r%t        d| j                   › d|j                   › �«      ‚| j                  «       |j                  «       k7  r-t        d| j                  «       › d|j                  «       › �«      ‚| |fS )z Check that auc input is correct.r6   zJExpected both `x` and `y` tensor to be 1d, but got tensors with dimension z and zHExpected the same number of elements in `x` and `y` tensor but received )ÚndimÚsqueezeÚ
ValueErrorÚnumelr   s     r   Ú_auc_format_inputsrB   `   sÈ   € à—v‘v ’zˆ�	‰	Œ q€AØ—v‘v ’zˆ�	‰	Œ q€Aà‡v�v�‚z�Q—V‘V˜a’ZÜØXÐYZ×Y_ÑY_ÐX`Ð`eÐfg×flÑflÐemÐnó
ð 	
ð 	‡w�wƒy�A—G‘G“IÒÜØVÐWX×W^ÑW^ÓW`ÐVaÐafÐgh×gnÑgnÓgpÐfqÐró
ð 	
ð ˆaˆ4€Kr   Ú	directionÚaxisc                 ó�   — t        j                  «       5  t        j                  || |¬«      |z  }ddd«       |S # 1 sw Y   S xY w)zrCompute area under the curve using the trapezoidal rule.

    Assumes increasing or decreasing order of `x`.

    ©ÚdimN)r   Úno_gradÚtrapz)r   r	   rC   rD   Ú	auc_scores        r   Ú_auc_compute_without_checkrK   p   s;   € ô 
�‰�Ü!ŸK™K¨¨1°$Ô7¸)ÑCˆ	÷ 
àÐ÷ 
àÐús	   •;»AÚreorderc                 ó4  — t        j                  «       5  |rt        j                  | d¬«      \  } }||   }| dd | dd z
  }|dk  j                  «       r!|dk  j	                  «       rd}nt        d«      ‚d	}t        | ||«      cddd«       S # 1 sw Y   yxY w)
zñCompute area under the curve using the trapezoidal rule.

    Example:
        >>> import torch
        >>> x = torch.tensor([1, 2, 3, 4])
        >>> y = torch.tensor([1, 2, 3, 4])
        >>> _auc_compute(x, y)
        tensor(7.5000)

    T)Ústabler6   Nr7   r   g      ð¿z_The `x` tensor is neither increasing or decreasing. Try setting the reorder argument to `True`.g      ð?)r   rH   Úsortr&   Úallr@   rK   )r   r	   rL   Úx_idxÚdxrC   s         r   Ú_auc_computerS   {   sŽ   € ô 
�‰�ÙÜ—z‘z !¨DÔ1‰HˆAˆuØ�%‘ˆAàˆqˆrˆU�Q�s˜�V‰^ˆØ�‰F�<‰<Œ>Ø�a‘�}‰}ŒØ ‘	ä Øuóð ð ˆIÜ)¨!¨Q°	Ó:÷ 
�Šús   •A/BÂBc                 ó<   — t        | |«      \  } }t        | ||¬«      S )a8  Compute Area Under the Curve (AUC) using the trapezoidal rule.

    Args:
        x: x-coordinates, must be either increasing or decreasing
        y: y-coordinates
        reorder: if True, will reorder the arrays to make it either increasing or decreasing

    Return:
        Tensor containing AUC score

    )rL   )rB   rS   )r   r	   rL   s      r   ÚaucrU   ˜   s#   € ô ˜a Ó#�D€A€qÜ˜˜1 gÔ.Ð.r   Úxpc           	      ó2  — t        |dd |dd z
  |dd |dd z
  «      }|dd ||dd z  z
  }t        j                  t        j                  | dd…df   |ddd…f   «      d«      dz
  }t        j                  |dt        |«      dz
  «      }||   | z  ||   z   S )ai  One-dimensional linear interpolation for monotonically increasing sample points.

    Returns the one-dimensional piecewise linear interpolation to a function with
    given discrete data points :math:`(xp, fp)`, evaluated at :math:`x`.

    Adjusted version of this https://github.com/pytorch/pytorch/issues/50334#issuecomment-1000917964

    Args:
        x: the :math:`x`-coordinates at which to evaluate the interpolated values.
        xp: the :math:`x`-coordinates of the data points, must be increasing.
        fp: the :math:`y`-coordinates of the data points, same length as `xp`.

    Returns:
        the interpolated values, same size as `x`.

    Example:
        >>> x = torch.tensor([0.5, 1.5, 2.5])
        >>> xp = torch.tensor([1, 2, 3])
        >>> fp = torch.tensor([1, 2, 3])
        >>> interp(x, xp, fp)
        tensor([0.5000, 1.5000, 2.5000])

    r6   Nr7   r   )r+   r   r:   ÚgeÚclampÚlen)r   rV   r0   ÚmÚbÚindicess         r   Úinterpr^   ¨   s­   € ô0 	�R˜˜�V˜b  "˜gÑ% r¨!¨" v°°3°B°Ñ'7Ó8€AØ
ˆ3ˆBˆ�1�r˜#˜2�w‘;Ñ€Aä�i‰iœŸ™ ¢1 d 7¡¨R°²a°©[Ó9¸1Ó=ÀÑA€GÜ�k‰k˜' 1¤c¨!£f¨q¡jÓ1€GàˆW‰:˜‰>˜A˜g™JÑ&Ð&r   ÚtensorÚnormalization)ÚsigmoidÚsoftmaxc                 ó®  — |s| S | j                   t        j                   d«      k(  rLt        j                  | dk\  | dk  z  «      s,|dk(  r| j                  «       nt        j                  | d¬«      } | S | dk  | dkD  z  j                  «       }t        j                  ||dk(  rt        j                  | «      | «      S t        j                  | d¬«      | «      S )aÚ  Normalize logits if needed.

    If input tensor is outside the [0,1] we assume that logits are provided and apply the normalization.
    Use torch.where to prevent device-host sync.

    Args:
        tensor: input tensor that may be logits or probabilities
        normalization: normalization method, either 'sigmoid' or 'softmax'

    Returns:
        normalized tensor if needed

    Example:
        >>> import torch
        >>> tensor = torch.tensor([-1.0, 0.0, 1.0])
        >>> normalize_logits_if_needed(tensor, normalization="sigmoid")
        tensor([0.2689, 0.5000, 0.7311])
        >>> tensor = torch.tensor([[-1.0, 0.0, 1.0], [1.0, 0.0, -1.0]])
        >>> normalize_logits_if_needed(tensor, normalization="softmax")
        tensor([[0.0900, 0.2447, 0.6652],
                [0.6652, 0.2447, 0.0900]])
        >>> tensor = torch.tensor([0.0, 0.5, 1.0])
        >>> normalize_logits_if_needed(tensor, normalization="sigmoid")
        tensor([0.0000, 0.5000, 1.0000])

    Úcpur   r6   ra   rF   )r"   r   rP   ra   rb   r&   r(   )r_   r`   Ú	conditions      r   Únormalize_logits_if_neededrf   É   sÄ   € ñ8 Øˆà‡}�}œŸ™ UÓ+Ò+Ü�y‰y˜& A™+¨&°A©+Ñ6Ô7Ø)6¸)Ò)C�V—^‘^Ô%ÌÏÉÐW]ÐcdÔIeˆFØˆð ˜1‘* ¨!¡Ñ,×1Ñ1Ó3€IÜ�;‰;ØØ!.°)Ò!;Œ�‰�fÓØóð äAFÇÁÈvÐ[\ÔA]Øóð r   )r   )r6   )r7   )F)Útypingr   r   r   r   Útyping_extensionsr   Útorchmetrics.utilitiesr   r   r   r   r+   ÚstrÚboolr%   r<   ÚtuplerB   rK   rS   rU   r^   rf   r!   r   r   Ú<module>rm      s·  ð÷ #ã Ý Ý %å 1ð�Fð ˜vð ¨&ó ð�6ð ˜fð ¨ó ð$ ;>ñ)Ø	ð)àð)ð ˜ ¨Ñ 6Ð6Ñ7ð)ð ó	)ðD opñPØðPØ$ S™MðPØ7;ðPØAGðPØMSðPØY_ðPØhkðPàóPð˜&ð  Vð °°f¸f°nÑ0Eó ñ  &ð ¨Vð Àð ÈSð ÐZ`ó ñ;�Fð ;˜vð ;°ð ;Àó ;ñ:/ˆ6ð /�fð / tð /¸ó /ð 'ˆfð '˜&ð ' fð '°ó 'ðB* vð *¸hÀwÐOcÑGdÑ>eð *Ðjpô *r   