Ë
    ýÿæio#  ã            	       ó   — d dl Z d dlmZ d dlmZ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mZ d dlmZ d	Zd
eeee   f   defd„Zd
edefd„Zd
edefd„Zd
edefd„Zd
edefd„Zd
edefd„Zd
edeeef   fd„Z	 d,dede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.d
ededefd„Z#d
edefd„Z$dedefd „Z%d,d
ed!ee   defd"„Z&d/d
edee   d#eejN                     defd$„Z(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+y)0é    N)ÚSequence)ÚAnyÚListÚOptionalÚUnion)Úapply_to_collection)ÚTensor)ÚTorchMetricsUserWarning)Ú_TORCH_LESS_THAN_2_6Ú_XLA_AVAILABLE)Úrank_zero_warng�íµ ÷Æ°>ÚxÚreturnc                 ó  — t        | t        j                  «      r| S | D �cg c]7  }|j                  «       dk(  r |j                  dk(  r|j                  d«      n|‘Œ9 } }| st        d«      ‚t        j                  | d¬«      S c c}w )z'Concatenation along the zero dimension.é   r   zNo samples to concatenate©Údim)Ú
isinstanceÚtorchr	   ÚnumelÚndimÚ	unsqueezeÚ
ValueErrorÚcat)r   Úys     úp/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchmetrics/utilities/data.pyÚdim_zero_catr      sr   € ä�!”U—\‘\Ô"ØˆÙJKÓLÉ!ÀQ˜1Ÿ7™7›9¨š>¨a¯f©f¸ªkˆ�‰�QŒ¸qÑ	@È!€AÐLÙÜÐ4Ó5Ð5Ü�9‰9�Q˜AÔÐùò 	Ms   ¡<Bc                 ó0   — t        j                  | d¬«      S )z#Summation along the zero dimension.r   r   )r   Úsum©r   s    r   Údim_zero_sumr!   '   s   € ä�9‰9�Q˜AÔÐó    c                 ó0   — t        j                  | d¬«      S )z!Average along the zero dimension.r   r   )r   Úmeanr    s    r   Údim_zero_meanr%   ,   s   € ä�:‰:�a˜QÔÐr"   c                 óD   — t        j                  | d¬«      j                  S )zMax along the zero dimension.r   r   )r   ÚmaxÚvaluesr    s    r   Údim_zero_maxr)   1   ó   € ä�9‰9�Q˜AÔ×%Ñ%Ð%r"   c                 óD   — t        j                  | d¬«      j                  S )zMin along the zero dimension.r   r   )r   Úminr(   r    s    r   Údim_zero_minr-   6   r*   r"   c                 ó@   — | D ��cg c]  }|D ]  }|‘Œ Œ c}}S c c}}w )z&Flatten list of list into single list.© )r   ÚsublistÚitems      r   Ú_flattenr2   ;   s"   € á !Ô6¡�W«g dŠD¨gˆD Ò6Ð6ùÓ6s   †c                 óÀ   — i }d}| j                  «       D ]D  \  }}t        |t        «      r$|j                  «       D ]  \  }}||v rd}|||<   Œ Œ:||v rd}|||<   ŒF ||fS )zYFlatten dict of dicts into single dict and checking for duplicates in keys along the way.FT)Úitemsr   Údict)r   Únew_dictÚ
duplicatesÚkeyÚvalueÚkÚvs          r   Ú_flatten_dictr<   @   sw   € à€HØ€JØ—g‘g–i‰
ˆˆUÜ�eœTÔ"ØŸ™ž‘��1Ø˜‘=Ø!%�JØ�˜’ñ &ð
 �h‰Ø!�
Ø!ˆH�SŠMð  ð �ZÐÐr"   Úlabel_tensorÚnum_classesc                 óŠ  — |€8t        | j                  «       j                  «       j                  «       dz   «      }t	        j
                  | j                  d   |g| j                  dd ¢­| j                  | j                  dœŽ}| j                  «       j                  d«      j                  |«      }|j                  d|d«      S )a¥  Convert  a dense label tensor to one-hot format.

    Args:
        label_tensor: dense label tensor, with shape [N, d1, d2, ...]
        num_classes: number of classes C

    Returns:
        A sparse label tensor with shape [N, C, d1, d2, ...]

    Example:
        >>> x = torch.tensor([1, 2, 3])
        >>> to_onehot(x)
        tensor([[0, 1, 0, 0],
                [0, 0, 1, 0],
                [0, 0, 0, 1]])

    Nr   r   )ÚdtypeÚdeviceç      ð?)Úintr'   Údetachr1   r   ÚzerosÚshaper@   rA   Úlongr   Ú	expand_asÚscatter_)r=   r>   Útensor_onehotÚindexs       r   Ú	to_onehotrL   Q   s¾   € ð* ÐÜ˜,×*Ñ*Ó,×3Ñ3Ó5×:Ñ:Ó<¸qÑ@ÓAˆä—K‘KØ×Ñ˜1ÑØðð 
×	Ñ	˜A˜BÐ	ñð × Ñ Ø×"Ñ"ò€Mð ×ÑÓ×)Ñ)¨!Ó,×6Ñ6°}ÓE€EØ×!Ñ! ! U¨CÓ0Ð0r"   r:   r   c                 ó  — | j                   t        j                  k(  rF| j                  s:t        j                  | |d¬«      j                  |«      }|j                  |d|«      S | j                  ||¬«      j                  S )z3torch.top_k does not support half precision on CPU.T)r   Ústabler   ©r:   r   )	r@   r   ÚhalfÚis_cudaÚargsortÚflipÚnarrowÚtopkÚindices)r   r:   r   Úidxs       r   Ú"_top_k_with_half_precision_supportrX   t   sa   € à‡w�w”%—*‘*Ò Q§Y¢YÜ�m‰m˜A 3¨tÔ4×9Ñ9¸#Ó>ˆØ�z‰z˜#˜q !Ó$Ð$Ø�6‰6�A˜3ˆ6Ó×'Ñ'Ð'r"   Úprob_tensorrU   c                 ó  — t        j                  | t         j                  ¬«      }|dk(  r4|j                  || j	                  |d¬«      d«       |j                  «       S |j                  |t        | ||¬«      d«       |j                  «       S )aw  Convert a probability tensor to binary by selecting top-k the highest entries.

    Args:
        prob_tensor: dense tensor of shape ``[..., C, ...]``, where ``C`` is in the
            position defined by the ``dim`` argument
        topk: number of the highest entries to turn into 1s
        dim: dimension on which to compare entries

    Returns:
        A binary tensor of the same shape as the input tensor of type ``torch.int32``

    Example:
        >>> x = torch.tensor([[1.1, 2.0, 3.0], [2.0, 1.0, 0.5]])
        >>> select_topk(x, topk=2)
        tensor([[0, 1, 1],
                [1, 1, 0]], dtype=torch.int32)

    ©r@   r   T)r   ÚkeepdimrB   rO   )r   Ú
zeros_likerC   rI   ÚargmaxrX   )rY   rU   r   Útopk_tensors       r   Úselect_topkr`   |   s€   € ô& ×"Ñ" ;´e·i±iÔ@€KØˆq‚yØ×Ñ˜S +×"4Ñ"4¸ÀdÐ"4Ó"KÈSÔQð �?‰?ÓÐð 	×Ñ˜SÔ"DÀ[ÐTXÐ^aÔ"bÐdgÔhØ�?‰?ÓÐr"   Ú
argmax_dimc                 ó0   — t        j                  | |¬«      S )aw  Convert  a tensor of probabilities to a dense label tensor.

    Args:
        x: probabilities to get the categorical label [N, d1, d2, ...]
        argmax_dim: dimension to apply

    Return:
        A tensor with categorical labels [N, d2, ...]

    Example:
        >>> x = torch.tensor([[0.2, 0.5], [0.9, 0.1]])
        >>> to_categorical(x)
        tensor([1, 0])

    r   )r   r^   )r   ra   s     r   Úto_categoricalrc   —   s   € ô  �<‰<˜˜zÔ*Ð*r"   c                 óL   — | j                  «       dk(  r| j                  «       S | S )Nr   )r   Úsqueezer    s    r   Ú_squeeze_scalar_element_tensorrf   ª   s   € ØŸ'™'›) qš.ˆ1�9‰9‹;Ð/¨aÐ/r"   Údatac                 ó,   — t        | t        t        «      S ©N)r   r	   rf   )rg   s    r   Ú_squeeze_if_scalarrj   ®   s   € Ü˜t¤VÔ-KÓLÐLr"   Ú	minlengthc                 óœ  — |€t        t        j                  | «      «      }t        j                  «       st        s| j
                  rpt        j                  || j                  ¬«      j                  t        | «      d«      }t        j                  | j                  dd«      |«      j                  d¬«      S t        j                  | |¬«      S )a   Implement custom bincount.

    PyTorch currently does not support ``torch.bincount`` when running in deterministic mode on GPU or when running
    MPS devices or when running on XLA device. This implementation therefore falls back to using a combination of
    `torch.arange` and `torch.eq` in these scenarios. A small performance hit can expected and higher memory consumption
    as `[batch_size, mincount]` tensor needs to be initialized compared to native ``torch.bincount``.

    Args:
        x: tensor to count
        minlength: minimum length to count

    Returns:
        Number of occurrences for each unique element in x

    Example:
        >>> x = torch.tensor([0,0,0,1,1,2,2,2,2])
        >>> _bincount(x, minlength=3)
        tensor([3, 2, 4])

    )rA   r   éÿÿÿÿr   r   ©rk   )Úlenr   ÚuniqueÚ$are_deterministic_algorithms_enabledr   Úis_mpsÚarangerA   ÚrepeatÚeqÚreshaper   Úbincount)r   rk   Úmeshs      r   Ú	_bincountry   ²   s�   € ð* ÐÜœŸ™ Q›Ó(ˆ	ä×1Ñ1Ô3µ~ÈÏÊÜ�|‰|˜I¨a¯h©hÔ7×>Ñ>¼sÀ1»vÀqÓIˆÜ�x‰x˜Ÿ	™	 " aÓ(¨$Ó/×3Ñ3¸Ð3Ó:Ð:ä�>‰>˜! yÔ1Ð1r"   r@   c                 ód  — t        j                  «       xr | j                  xr | j                  «       }t        r_|r]t
        j                  dk7  rJt        dt        «       | j                  «       j                  ||¬«      j                  | j                  «      S t        j                  | ||¬«      S )z\Implement custom cumulative summation for Torch versions which does not support it natively.Úwin32zåYou are trying to use a metric in deterministic mode on GPU that uses `torch.cumsum`, which is currently not supported. The tensor will be copied to the CPU memory to compute it and then copied back to GPU. Expect some slowdowns.)r   r@   )r   rq   rQ   Úis_floating_pointr   ÚsysÚplatformr   r
   ÚcpuÚcumsumÚtorA   )r   r   r@   Úis_cuda_fp_deterministics       r   Ú_cumsumrƒ   Ñ   sŒ   € ä$×IÑIÓKÒsÐPQ×PYÑPYÒsÐ^_×^qÑ^qÓ^sÐÝÑ 8¼S¿\¹\ÈWÒ=TÜð&ô $ô		
ð �u‰u‹w�~‰~ #¨Uˆ~Ó3×6Ñ6°q·x±xÓ@Ð@Ü�<‰<˜˜s¨%Ô0Ð0r"   c                 ób   — t        j                  | d¬«      \  }}t        |t        |«      ¬«      S )zÎSimilar to `_bincount`, but works also with tensor that do not contain continuous values.

    Args:
        x: tensor to count

    Returns:
        Number of occurrences for each unique element in x

    T)Úreturn_inversern   )r   rp   ry   ro   )r   Úunique_xÚinverse_indicess      r   Ú_flexible_bincountrˆ   ß   s*   € ô !&§¡¨Q¸tÔ DÑ€HˆoÜ�_´°H³Ô>Ð>r"   Útensor1Útensor2c                 ó˜   — | j                   |j                   k7  r|j                  | j                   ¬«      }t        j                  | |«      S )z:Wrap torch.allclose to be robust towards dtype difference.r[   )r@   r�   r   Úallclose)r‰   rŠ   s     r   rŒ   rŒ   í   s7   € à‡}�}˜Ÿ™Ò%Ø—*‘* 7§=¡=�*Ó1ˆÜ�>‰>˜' 7Ó+Ð+r"   ÚxpÚfpc                 ó  — t        j                  |«      }||   }||   }|dd |dd z
  |dd |dd z
  z  }t        j                  || «      dz
  }t        j                  |dt	        |«      dz
  «      }||   ||   | ||   z
  z  z   S )zàInterpolation function comparable to numpy.interp.

    Args:
        x: x-coordinates where to evaluate the interpolated values
        xp: x-coordinates of the data points
        fp: y-coordinates of the data points

    r   Nrm   r   )r   rR   ÚsearchsortedÚclampro   )r   r�   rŽ   Úsorted_indicesÚslopesrV   s         r   Úinterpr”   ô   s©   € ô —]‘] 2Ó&€NØ	ˆNÑ	€BØ	ˆNÑ	€Bð ��ˆf�r˜#˜2�wÑ 2 a b 6¨B¨s°¨GÑ#3Ñ4€Fô × Ñ   QÓ'¨!Ñ+€GÜ�k‰k˜' 1¤c¨&£k°A¡oÓ6€Gð ˆg‰;˜ ™¨A°°7±©OÑ<Ñ<Ð<r"   ri   )r   r   )r   )r   N),r}   Úcollections.abcr   Útypingr   r   r   r   r   Úlightning_utilitiesr   r	   Ú!torchmetrics.utilities.exceptionsr
   Útorchmetrics.utilities.importsr   r   Útorchmetrics.utilities.printsr   Ú
METRIC_EPSr   r!   r%   r)   r-   Úlistr2   r5   ÚtupleÚboolr<   rC   rL   rX   r`   rc   rf   rj   ry   r@   rƒ   rˆ   rŒ   r”   r/   r"   r   Ú<module>rŸ      s  ðó Ý $ß -Ó -ã Ý 3Ý å Eß OÝ 8à€
ð�E˜& $ v¡,Ð.Ñ/ð °Fó ð�Fð ˜vó ð
 �Vð   ó  ð
&�Fð &˜vó &ð
&�Fð &˜vó &ð
7�ð 7˜Tó 7ð
 �Tð  ˜e D¨$ JÑ/ó  ð& "&ñ 1Øð 1à˜#‘ð 1ð ó 1ñF(¨&ð (°Sð (À3ð (Èvó (ñ˜Vð ¨3ð ¸ð ÀVó ñ6+�fð +¨#ð +°fó +ð&0 fð 0°ó 0ðM˜Sð M Só Mñ2�ð 2 H¨S¡Mð 2¸Vó 2ñ>1ˆvð 1˜H S™Mð 1°h¸u¿{¹{Ñ6Kð 1ÐW]ó 1ð?˜&ð ? Vó ?ð,�fð , vð ,°$ó ,ð=ˆfð =˜&ð = fð =°ô =r"   