You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch中高效计算张量内向量对欧氏距离/点积的最优方案

PyTorch高效计算百万组向量对的距离并定位最近对

针对你手里shape为(4000000,3,90)的Tensor(400万组,每组3个90维向量),要计算每组内所有向量对的欧氏距离(或点积)并找出每组最接近的向量对,绝对不能用循环——Python循环处理百万级数据效率极低,完全向量化的PyTorch内置操作才是最优解,以下是具体实现和最佳实践:

一、欧氏距离计算方案

方案1:手动计算向量差值的距离平方(性能最优)

因为只需要找最近对,不需要精确的距离值,用距离平方代替欧氏距离可以省去开根号的开销,速度更快:

import torch

# 必须移至GPU,CPU处理400万行会极慢
x = torch.randn(4000000, 3, 90).cuda()

# 计算每组内三对向量的差值
diff_01 = x[:, 0] - x[:, 1]
diff_02 = x[:, 0] - x[:, 2]
diff_12 = x[:, 1] - x[:, 2]

# 计算每对的距离平方
dist_sq_01 = torch.sum(diff_01 ** 2, dim=1)
dist_sq_02 = torch.sum(diff_02 ** 2, dim=1)
dist_sq_12 = torch.sum(diff_12 ** 2, dim=1)

# 合并为shape (4000000, 3)的距离矩阵
dist_sq_matrix = torch.stack([dist_sq_01, dist_sq_02, dist_sq_12], dim=1)

# 找出每组最小距离平方和对应的向量对索引
# 索引0对应(0,1),1对应(0,2),2对应(1,2)
min_dist_sq, min_pair_idx = torch.min(dist_sq_matrix, dim=1)

# 如果需要实际欧氏距离,再开根号
# min_dist = torch.sqrt(min_dist_sq)

方案2:用PyTorch内置torch.cdist(代码最简洁)

torch.cdist是高度优化的距离计算函数,能快速计算两组向量的两两距离,适合追求代码简洁的场景:

import torch

x = torch.randn(4000000, 3, 90).cuda()

# 计算每组内所有向量对的欧氏距离矩阵,shape (4000000,3,3)
dist_matrix = torch.cdist(x, x, p=2)

# 提取我们需要的三对非对角线距离
distances = torch.stack([
    dist_matrix[:, 0, 1],
    dist_matrix[:, 0, 2],
    dist_matrix[:, 1, 2]
], dim=1)

# 找最近对
min_dist, min_pair_idx = torch.min(distances, dim=1)

注:cdist会计算所有9对距离(包括向量自身),但因为内置实现高度优化,性能和手动计算差异很小,胜在代码简洁。

二、用点积衡量相似性的方案

如果用点积表示向量相似性(点积越大,向量越相似,归一化后等价于余弦相似度),实现逻辑类似:

import torch

x = torch.randn(4000000, 3, 90).cuda()

# 计算每组内三对向量的点积
dot_01 = torch.sum(x[:, 0] * x[:, 1], dim=1)
dot_02 = torch.sum(x[:, 0] * x[:, 2], dim=1)
dot_12 = torch.sum(x[:, 1] * x[:, 2], dim=1)

# 合并点积矩阵
dot_matrix = torch.stack([dot_01, dot_02, dot_12], dim=1)

# 点积越大越相似,所以取最大值对应的索引
max_dot, max_pair_idx = torch.max(dot_matrix, dim=1)

三、最佳实践

  1. 强制用GPU:400万行的计算在CPU上会耗时数分钟甚至更久,移至CUDA设备后能把时间压缩到秒级。
  2. 优先用距离平方:如果不需要精确距离值,用距离平方代替欧氏距离,省去开根号的计算开销。
  3. 内存不够时批量处理:如果GPU显存较小(比如<8GB),可以把数据分成批量处理,避免OOM:
batch_size = 100000
min_dist_sq_list = []
min_pair_idx_list = []

for i in range(0, x.shape[0], batch_size):
    batch_x = x[i:i+batch_size]
    # 复用前面的距离平方计算逻辑
    diff_01 = batch_x[:,0] - batch_x[:,1]
    diff_02 = batch_x[:,0] - batch_x[:,2]
    diff_12 = batch_x[:,1] - batch_x[:,2]
    dist_sq_01 = torch.sum(diff_01**2, dim=1)
    dist_sq_02 = torch.sum(diff_02**2, dim=1)
    dist_sq_12 = torch.sum(diff_12**2, dim=1)
    dist_sq_matrix = torch.stack([dist_sq_01, dist_sq_02, dist_sq_12], dim=1)
    batch_min_dist_sq, batch_min_pair_idx = torch.min(dist_sq_matrix, dim=1)
    min_dist_sq_list.append(batch_min_dist_sq)
    min_pair_idx_list.append(batch_min_pair_idx)

# 合并所有批量结果
min_dist_sq = torch.cat(min_dist_sq_list)
min_pair_idx = torch.cat(min_pair_idx_list)
  1. 避免显式循环:所有操作都用PyTorch的向量化API,Python循环在百万级数据面前效率可以忽略不计。

内容的提问来源于stack exchange,提问作者WalaWizon

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.01 22:41:12