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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 13:58:25