PyTorch torch.cdist Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.cdistIt is a function in PyTorch used to compute the Euclidean distance matrix between two sets of points. It calculates the Euclidean distance between each point in the first input and each point in the second input, returning a distance matrix.

Function Definition

torch.cdist(input1, input2, p=2.0, compute_mode='use_mm_for_euclidean_dist')

Usage Example

Example

import torch

# Compute the Euclidean distance between two sets of points
x = torch.tensor([[0, 0], [1, 1], [2, 2]])  # 3 2D points
y = torch.tensor([[0, 0], [1, 0], [2, 0]])  # 3 2D points

# Distance matrix shape: (3, 3)
distances = torch.cdist(x, y)
print("Point x:")
print(x)
print("Point y:")
print(y)
print("Euclidean distance matrix:")
print(distances)

# Use a different p value (Manhattan distance p=1)
dist_l1 = torch.cdist(x, y, p=1.0)
print("L1 distance (p=1):")
print(dist_l1)

# Use p=infinity (Chebyshev distance)
dist_inf = torch.cdist(x, y, p=float('inf'))
print("Chebyshev distance (p=inf):")
print(dist_inf)

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions