Функция cdist
Функция cdist вычисляет попарные расстояния между каждым вектором в первом наборе и каждым вектором во втором наборе.
Первый параметр x1 - это тензор размерности (B, P, M) или (P, M),
где P - количество векторов, M - размерность каждого вектора.
Второй параметр x2 - тензор размерности (B, R, M) или (R, M),
где R - количество векторов.
Третий параметр p - показатель степени для вычисления расстояния Минковского (по умолчанию 2 для евклидова расстояния).
Четвертый параметр compute_mode - режим вычислений: 'use_mm_for_euclid_dist' (использовать перемножение матриц для ускорения), 'donot_use_mm_for_euclid_dist' (не использовать перемножение матриц) или None (автоматический выбор).
Функция возвращает тензор расстояний размерности (B, P, R) или (P, R).
Синтаксис
torch.cdist(x1, x2, p=2, compute_mode=None)
Пример
Вычислим евклидовы расстояния между векторами из двух наборов:
import torch
x1 = torch.tensor([
[1.0, 2.0],
[3.0, 4.0],
])
x2 = torch.tensor([
[5.0, 6.0],
[7.0, 8.0],
])
res = torch.cdist(x1, x2)
print(res)
Результат выполнения кода:
tensor([
[5.6569, 8.4853],
[2.8284, 5.6569],
])
Пример
Используем расстояние Минковского с p=1 (манхэттенское расстояние):
import torch
x1 = torch.tensor([
[1.0, 2.0],
[3.0, 4.0],
])
x2 = torch.tensor([
[5.0, 6.0],
[7.0, 8.0],
])
res = torch.cdist(x1, x2, p=1)
print(res)
Результат выполнения кода:
tensor([
[8., 12.],
[4., 8.],
])
Пример
Работа с пакетной размерностью (два набора по два вектора):
import torch
x1 = torch.tensor([
[
[1.0, 2.0],
[3.0, 4.0],
],
[
[5.0, 6.0],
[7.0, 8.0],
],
])
x2 = torch.tensor([
[
[9.0, 10.0],
[11.0, 12.0],
],
[
[13.0, 14.0],
[15.0, 16.0],
],
])
res = torch.cdist(x1, x2)
print(res.shape)
print(res)
Результат выполнения кода:
torch.Size([2, 2, 2])
tensor([
[
[11.3137, 14.1421],
[7.0711, 9.8995],
],
[
[11.3137, 14.1421],
[7.0711, 9.8995],
],
])
Пример
Вычисление расстояний с фиксацией генератора случайных чисел для воспроизводимости:
import torch
torch.manual_seed(0)
x1 = torch.randn(3, 5)
x2 = torch.randn(4, 5)
res = torch.cdist(x1, x2)
print(res.shape)
print(res)
Результат выполнения кода:
torch.Size([3, 4])
tensor([
[2.4616, 2.6200, 2.1121, 2.5824],
[2.1258, 1.8386, 2.2587, 1.2680],
[2.6824, 3.0084, 2.3330, 2.8779],
])