如何提升3D网络节点锥状求和for循环的计算速度?
优化方案
1. 预分组Z层节点
原代码每次遍历整个network查找对应Z层节点,这是核心性能瓶颈。先按Z坐标对节点分组,后续直接调用对应层数据:
from collections import defaultdict # 预处理:按z值分组节点,存储(x, y, value) z_groups = defaultdict(lambda: {'x': [], 'y': [], 'val': []}) for node in network: x, y, z, val = node z_groups[z]['x'].append(x) z_groups[z]['y'].append(y) z_groups[z]['val'].append(val) # 转换为numpy数组进一步提速 import numpy as np for z in z_groups: z_groups[z]['x'] = np.array(z_groups[z]['x']) z_groups[z]['y'] = np.array(z_groups[z]['y']) z_groups[z]['val'] = np.array(z_groups[z]['val'])
2. 用Numpy向量化筛选求和
基于分组后的Z层数据,用numpy布尔索引替代Python循环筛选节点,求和速度提升显著:
import numpy as np for item in network: item_x, item_y, item_z, _ = item pw = 0 # 避免z-level小于0,限制最大遍历层数 max_level = min(item_z - 1, rl_alternatives - 1) for level in range(1, max_level + 1): target_z = item_z - level if target_z not in z_groups: continue x_arr = z_groups[target_z]['x'] y_arr = z_groups[target_z]['y'] val_arr = z_groups[target_z]['val'] # 布尔索引快速筛选符合范围的节点 mask = (x_arr >= item_x - level) & (x_arr <= item_x + level) & \ (y_arr >= item_y - level) & (y_arr <= item_y + level) pw += val_arr[mask].sum() item.append(pw)
3. 为Z层构建二维前缀和(适合坐标范围固定场景)
若x、y坐标范围连续且已知,可为每个Z层构建前缀和数组,实现O(1)时间查询矩形区域和:
# 预处理每个Z层的前缀和 z_prefix_sums = {} min_x = min(node[0] for node in network) max_x = max(node[0] for node in network) min_y = min(node[1] for node in network) max_y = max(node[1] for node in network) x_offset = -min_x # 让x坐标从0开始 y_offset = -min_y for z in z_groups: # 创建对应尺寸的网格,初始值为0 grid = np.zeros((max_x - min_x + 1, max_y - min_y + 1), dtype=np.float64) x_list = z_groups[z]['x'] y_list = z_groups[z]['y'] val_list = z_groups[z]['val'] # 将节点值填充到网格对应位置 for x, y, val in zip(x_list, y_list, val_list): grid[x + x_offset][y + y_offset] += val # 计算二维前缀和 prefix_sum = np.cumsum(np.cumsum(grid, axis=0), axis=1) z_prefix_sums[z] = prefix_sum # 矩形区域和查询函数 def get_rect_sum(prefix_sum, x1, y1, x2, y2): # 转换为网格坐标 x1 += x_offset y1 += y_offset x2 += x_offset y2 += y_offset # 处理边界超出情况 x1 = max(x1, 0) y1 = max(y1, 0) x2 = min(x2, prefix_sum.shape[0]-1) y2 = min(y2, prefix_sum.shape[1]-1) if x1 > x2 or y1 > y2: return 0 # 前缀和公式计算区域和 total = prefix_sum[x2][y2] if x1 > 0: total -= prefix_sum[x1-1][y2] if y1 > 0: total -= prefix_sum[x2][y1-1] if x1 > 0 and y1 > 0: total += prefix_sum[x1-1][y1-1] return total # 计算每个节点的位置值 for item in network: item_x, item_y, item_z, _ = item pw = 0 max_level = min(item_z - 1, rl_alternatives - 1) for level in range(1, max_level + 1): target_z = item_z - level if target_z not in z_prefix_sums: continue prefix_sum = z_prefix_sums[target_z] x1 = item_x - level y1 = item_y - level x2 = item_x + level y2 = item_y + level pw += get_rect_sum(prefix_sum, x1, y1, x2, y2) item.append(pw)
4. 批量处理减少重复计算
若存在大量节点需要计算,可批量处理每个Z层对所有符合条件节点的贡献:遍历每个Z层z,计算level = item_z - z对应的节点,一次性完成该层对所有相关节点的贡献累加,避免每个节点单独遍历Z层。
原优化效果差的原因
- Cython仅提速10%:原代码瓶颈是遍历整个network的冗余操作,而非循环计算本身,Cython无法抵消遍历大量节点的开销。
- Multiprocessing变慢:进程间通信的开销远大于并行计算的收益,尤其是单任务计算量较小时,进程启动和数据传递成本更高。
内容的提问来源于stack exchange,提问作者KHas
相关产品推荐
相关产品推荐

