如何优化numpy场景下计算坐标两两距离的循环效率?
代码优化方案
你当前使用的纯Python双层循环时间复杂度为O(n²),39000条数据对应15亿次以上的运算,Python原生循环的解释器开销会导致运行速度极慢,推荐以下两个优化方向:
方案1:使用SciPy向量化计算(最快,适合标准经纬度距离场景)
经纬度距离计算一般用Haversine球面距离公式,SciPy的cdist函数底层为C实现,支持直接批量计算所有点对的距离,完全避免Python循环开销,速度可以提升上千倍。
import numpy as np from scipy.spatial.distance import cdist # 提取id和经纬度,经纬度转弧度(Haversine公式要求输入为弧度) ids = a[:, 0].astype(int) lat_lon_rad = np.radians(a[:, 1:]) # 批量计算所有点对的距离,乘以地球半径6371得到单位为km的结果 dist_matrix = cdist(lat_lon_rad, lat_lon_rad, metric='haversine') * 6371 # 构造符合要求的输出数组 n = len(ids) i_col = np.repeat(ids, n) j_col = np.tile(ids, n) dist_col = dist_matrix.flatten() # 过滤掉点和自身配对的记录 mask = i_col != j_col result = np.column_stack([i_col[mask], j_col[mask], dist_col[mask]])
方案2:使用Numba JIT编译(改动最小,适合自定义距离函数)
如果你使用的getDistance是自定义逻辑,无法用SciPy内置度量实现,可以用Numba对循环做即时编译,编译后运行速度接近C语言,改动成本极低。
import numpy as np from numba import jit # 给自定义距离函数加JIT装饰器 @jit(nopython=True) def getDistance(lat1, lon1, lat2, lon2): # 保留你原有距离计算逻辑不变 pass # 给主循环加JIT装饰器 @jit(nopython=True) def calc_all_distances(a): n = a.shape[0] # 提前分配所有结果的内存,避免动态扩容开销 result = np.zeros((n * (n - 1), 3)) line = 0 for i in range(n): id_i, lat_i, lon_i = a[i] for j in range(n): if i == j: continue id_j, lat_j, lon_j = a[j] dis = getDistance(lat_i, lon_i, lat_j, lon_j) result[line] = [id_i, id_j, dis] line += 1 return result # 调用函数得到结果 result = calc_all_distances(a)
注意事项
- 39000条数据的非自身配对总共有约15.2亿条记录,单条记录如果存3个8字节数值,总内存占用约36GB,如果你的设备内存不足,建议分块计算,每次取部分点和全量点计算后写入硬盘,再处理下一批。
- 如果允许调整计算逻辑,也可以先只计算上三角的点对(m<n),再复制一份交换id顺序得到(n,m)的配对,能节省一半计算量。
内容的提问来源于stack exchange,提问作者Manfred
相关产品推荐
相关产品推荐

