PyTorch中平方距离计算:如何避免使用for循环
问题描述
需要计算尺寸为(20,20)的2D网格展平后(共400个点),所有网格点与指定点集的平方距离。当前通过for循环实现,希望去掉循环直接得到成对平方距离矩阵best_distance_squares。
原代码如下:
# best locations of indices existing on 2D grid- best_loc.shape # torch.Size([1024, 2]) # Specify 2D grid size- m = 20 n = 20 locs = [np.array([i, j]) for i in range(m) for j in range(n)] locations = torch.LongTensor(np.array(locs)) locations.shape # torch.Size([400, 2]) def get_distance_squares(best_loc): ''' Compute squared distances between 'best_loc' and 'locations' ''' best_loc = best_loc.unsqueeze(0).expand_as(locations).float() best_distance_squares = torch.sum(torch.pow(locations.float() - best_loc, 2), 1) return best_distance_squares bmu_distance_squares = list() for loc in bmu_loc: bmu_distance_squares.append(get_distance_squares(loc)) best_distance_squares = torch.stack(best_distance_squares) best_distance_squares.shape # torch.Size([1024, 400])
解决方案
可以利用PyTorch的广播机制或者平方距离的矩阵运算公式实现无循环计算,后者效率更高,适合大规模数据场景。
方法1:广播机制直接扩展维度
通过扩展best_loc和locations的维度,让PyTorch自动广播完成逐元素运算,最后求和得到平方距离:
# 先转换为float类型,避免整数运算精度问题 best_loc_float = best_loc.float() locations_float = locations.float() # 扩展维度:best_loc -> (1024, 1, 2),locations -> (1, 400, 2) # 广播后逐元素相减、平方,最后在最后一维求和 best_distance_squares = torch.sum( torch.pow(best_loc_float.unsqueeze(1) - locations_float.unsqueeze(0), 2), dim=2 ) # 结果形状为torch.Size([1024, 400]),与原代码输出一致
方法2:利用平方距离数学公式(更高效)
平方距离可拆解为数学公式:
$$||a - b||^2 = ||a||^2 + ||b||^2 - 2a \cdot b$$
通过矩阵乘法实现点积运算,避免逐元素操作,运算速度更快:
best_loc_float = best_loc.float() locations_float = locations.float() # 计算每个点的平方范数 best_norm = torch.sum(best_loc_float ** 2, dim=1, keepdim=True) # shape (1024, 1) locations_norm = torch.sum(locations_float ** 2, dim=1) # shape (400,) # 计算点积矩阵:(1024,2) @ (2,400) = (1024,400) dot_product = best_loc_float @ locations_float.T # 套用公式计算平方距离 best_distance_squares = best_norm + locations_norm - 2 * dot_product # 结果形状同样为torch.Size([1024, 400])
说明
- 两种方法均完全规避for循环,利用PyTorch的向量化运算提升效率,其中方法2的时间复杂度更低,更适合处理更大规模的点集。
- 必须先将张量转为
float类型,避免整数运算导致的精度丢失或错误。
内容的提问来源于stack exchange,提问作者Arun
相关产品推荐
相关产品推荐

