Ë
    þÍ:jy  ã                   óž   — d dl mZ d dlZd dlmZ d dlmZ d dlmZmZ 	 ddedee   dee	   d	efd
„Z
	 	 	 ddedee   ded   dee	   d	ef
d„Zy)é    )ÚOptionalN)ÚTensor)ÚLiteral)Ú_check_inputÚ_reduce_distance_matrixÚxÚyÚzero_diagonalÚreturnc                 óº  — t        | ||«      \  } }}| j                  }| j                  t        j                  «      } |j                  t        j                  «      }| | z  j                  dd¬«      }||z  j                  d¬«      }||z   d| j                  |j                  «      z  z
  j                  |«      }|r|j                  d«       |j                  «       S )zëCalculate the pairwise euclidean distance matrix.

    Args:
        x: tensor of shape ``[N,d]``
        y: tensor of shape ``[M,d]``
        zero_diagonal: determines if the diagonal of the distance matrix should be set to zero

    é   T)ÚdimÚkeepdim)r   é   r   )
r   ÚdtypeÚtoÚtorchÚfloat64ÚsumÚmmÚTÚfill_diagonal_Úsqrt)r   r	   r
   Ú_orig_dtypeÚx_normÚy_normÚdistances          ú/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/pairwise/euclidean.pyÚ#_pairwise_euclidean_distance_updater      s¶   € ô ' q¨!¨]Ó;Ñ€A€qˆ-à—'‘'€KØ	�‰ŒU�]‰]Ó€AØ	�‰ŒU�]‰]Ó€AØ�!‰e�[‰[˜Q¨ˆ[Ó-€FØ�!‰e�[‰[˜Qˆ[Ó€FØ˜‘ ! a§d¡d¨1¯3©3£i¡-Ñ/×3Ñ3°KÓ@€HÙØ×Ñ Ô"Ø�=‰=‹?Ðó    Ú	reduction)Úmeanr   ÚnoneNc                 ó4   — t        | ||«      }t        ||«      S )aþ  Calculate pairwise euclidean distances.

    .. math::
        d_{euc}(x,y) = ||x - y||_2 = \sqrt{\sum_{d=1}^D (x_d - y_d)^2}

    If both :math:`x` and :math:`y` are passed in, the calculation will be performed pairwise between
    the rows of :math:`x` and :math:`y`.
    If only :math:`x` is passed in, the calculation will be performed between the rows of :math:`x`.

    Args:
        x: Tensor with shape ``[N, d]``
        y: Tensor with shape ``[M, d]``, optional
        reduction: reduction to apply along the last dimension. Choose between `'mean'`, `'sum'`
            (applied along column dimension) or  `'none'`, `None` for no reduction
        zero_diagonal: if the diagonal of the distance matrix should be set to 0. If only `x` is given
            this defaults to `True` else if `y` is also given it defaults to `False`

    Returns:
        A ``[N,N]`` matrix of distances if only ``x`` is given, else a ``[N,M]`` matrix

    Example:
        >>> import torch
        >>> from torchmetrics.functional.pairwise import pairwise_euclidean_distance
        >>> x = torch.tensor([[2, 3], [3, 5], [5, 8]], dtype=torch.float32)
        >>> y = torch.tensor([[1, 0], [2, 1]], dtype=torch.float32)
        >>> pairwise_euclidean_distance(x, y)
        tensor([[3.1623, 2.0000],
                [5.3852, 4.1231],
                [8.9443, 7.6158]])
        >>> pairwise_euclidean_distance(x)
        tensor([[0.0000, 2.2361, 5.8310],
                [2.2361, 0.0000, 3.6056],
                [5.8310, 3.6056, 0.0000]])

    )r   r   )r   r	   r!   r
   r   s        r   Úpairwise_euclidean_distancer%   /   s    € ôR 3°1°a¸ÓG€HÜ" 8¨YÓ7Ð7r    )NN)NNN)Útypingr   r   r   Útyping_extensionsr   Ú(torchmetrics.functional.pairwise.helpersr   r   Úboolr   r%   © r    r   ú<module>r+      s˜   ðõ ã Ý Ý %ç Zð LPñØðØ˜6Ñ"ðØ:BÀ4¹.ðàóð4 Ø6:Ø$(ñ	*8Øð*8à�Ñð*8ð Ð2Ñ3ð*8ð ˜D‘>ð	*8ð
 ô*8r    