不同形状张量间的最小与平均欧氏距离计算优化求助
解决方案:高效计算欧氏距离并避免内存溢出
你的问题核心在于Python循环的开销和不必要的大尺寸临时张量生成——每次循环中A[row_id, :] - B会创建[100000, 14]的张量,累积占用显存且未利用GPU并行能力。以下是基于欧氏距离数学展开式的向量化优化方案,同时兼顾内存效率和计算速度。
核心原理:欧氏距离的矩阵展开
欧氏距离的平方可拆解为:
$$|a - b|^2 = |a|^2 + |b|^2 - 2a \cdot b$$
通过这个公式,我们可以用矩阵运算直接生成所有行对的距离平方矩阵,避免生成[1000, 100000, 14]的超大中间张量。
完整实现代码
import torch # 假设A、B已加载为PyTorch张量,示例: # A = torch.randn(1000, 14) # B = torch.randn(100000, 14) # 切换到GPU(如果可用) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") A = A.to(device) B = B.to(device) # 计算距离平方矩阵的基础分量 norm_A = torch.sum(A ** 2, dim=1, keepdim=True) # 形状 [1000, 1] norm_B = torch.sum(B ** 2, dim=1, keepdim=True) # 形状 [100000, 1] dot_product = A @ B.T # 形状 [1000, 100000],A与B的点积矩阵 # 计算欧氏距离平方,修正浮点精度导致的极小负数 dist_sq = norm_A + norm_B.T - 2 * dot_product dist_sq = dist_sq.clamp(min=0.0)
1. 计算每行最小距离的平均值
# 先取距离平方的最小值再开根号(比先开根号再取min更高效) min_dist_sq = dist_sq.min(dim=1)[0] min_dist = torch.sqrt(min_dist_sq) result_min_avg = min_dist.mean().item()
2. 计算每行平均距离的平均值
# 先计算所有欧氏距离,再按行取平均,最后求整体平均 dist = torch.sqrt(dist_sq) row_avg_dist = dist.mean(dim=1) result_avg_avg = row_avg_dist.mean().item()
方案优势
- 内存效率:生成的
dist_sq张量为[1000, 100000],float32格式下仅占约400MB,远低于T4 8GB显存上限,完全避免溢出。 - 计算速度:向量化矩阵运算充分利用GPU并行能力,比原循环方案快几十到上百倍,CPU运行效率也会显著提升。
超大规模B的备选分块方案
如果B的行数远超10万(如百万级),可以分块处理B,进一步降低显存占用:
block_size = 20000 # 根据显存调整块大小 min_dist_list = [] for i in range(0, B.shape[0], block_size): B_block = B[i:i+block_size] norm_B_block = torch.sum(B_block**2, dim=1, keepdim=True) dot_block = A @ B_block.T dist_sq_block = norm_A + norm_B_block.T - 2 * dot_block dist_sq_block = dist_sq_block.clamp(min=0.0) min_dist_block = torch.sqrt(dist_sq_block.min(dim=1)[0]) min_dist_list.append(min_dist_block) min_dist = torch.cat(min_dist_list) result_min_avg = min_dist.mean().item()
内容的提问来源于stack exchange,提问作者Tanay Rastogi
相关产品推荐
相关产品推荐

