NumPy矩阵基于索引数组的高效条件广播赋值实现
NumPy指定位置子矩阵条件赋值高效实现
核心实现思路
针对无循环批量赋值的需求,全程使用NumPy向量化操作实现,避免Python层循环带来的性能损耗:
- 用
np.ix_生成目标行列的索引网格,直接定位原矩阵中需要被赋值的矩形块 - 向量化判断两类赋值条件:原位置非元组、子矩阵元组第二个元素小于原位置元组第二个元素
- 用布尔掩码直接完成批量赋值,所有运算走NumPy底层C实现,性能远高于手动遍历
代码实现
import numpy as np # ---------- 示例数据 ---------- nparray = np.array([[0., 0. , 0., 0., 0., 0., 0., 0., 0., 0.], [0., (2, 2.5), 0., 0., 0., 0., 0., 0., 0., 0.], [0., (1, 6.5), 0., 0., 0., 0., 0., 0., 0., 0.], [0., 0. , 0., 0., 0., 0., 0., 0., 0., 0.], [0., 0. , 0., 0., 0., 0., 0., 0., 0., 0.], [0., 0. , 0., 0., 0., 0., 0., 0., 0., 0.], [0., 0. , 0., 0., 0., 0., 0., 0., 0., 0.], [0., 0. , 0., 0., 0., 0., 0., 0., 0., 0.], [0., 0. , 0., 0., 0., 0., 0., 0., 0., 0.], [0., 0. , 0., 0., 0., 0., 0., 0., 0., 0.]], dtype=object) sub_array= np.array([[(1, 3.2) , (2, 3.2), (3, 4.6), (4, 3.4)], [(3, 4.5) , (4, 0.4), (5, 3.2), (6, 2.3)], [(3, 4.5) , (5, 2.3), (7, 5.3), (9, 2.3)], [(12, 3.2), (45, 2.4), (32, 2.3), (6, 5.4)]], dtype=object) index = [1, 2, 5, 9] # ---------- 示例数据结束 ---------- # 核心赋值逻辑 # 1. 生成目标区域的二维索引网格,切出原矩阵对应块 rows, cols = np.ix_(index, index) target_block = nparray[rows, cols] # 2. 向量化判断原块每个位置是否为元组 is_tuple = np.vectorize(lambda x: isinstance(x, tuple))(target_block) # 3. 提取原位置和子矩阵的比较值(元组第二个元素),非元组位置原比较值设为无穷大 ori_val = np.where(is_tuple, np.vectorize(lambda x:x[1])(target_block), np.inf) sub_val = np.vectorize(lambda x:x[1])(sub_array) # 4. 生成赋值掩码:非元组直接赋值,是元组则子矩阵值更小才赋值 mask = (~is_tuple) | (sub_val < ori_val) # 5. 批量完成赋值 nparray[rows[mask], cols[mask]] = sub_array[mask]
运行后结果和预期完全一致:(1,1)位置原值2.5小于子矩阵值3.2,不替换;(2,1)位置原值6.5大于子矩阵值4.5,正常替换。
大规模数据优化建议
当前用object类型存储(索引, 值)元组的方式内存开销极高,11万规模的方阵用object存储内存占用会达到TB级,根本无法加载进内存,建议做如下调整:
- 拆分存储为两个独立的数值型矩阵:一个存距离值(初始填充无穷大,用float32类型),一个存邻点索引(用int32类型),总内存占用可以压缩到100GB以内,如果配合分块计算、稀疏矩阵存储还能进一步降低内存需求
- 生成
sub_array时直接拆分为索引数组和值数组,不要存成元组的object数组,省掉后续提取值的开销,速度还能再提升一个量级
拆分存储后的赋值逻辑更简洁,速度更快:
# 初始化矩阵示例(11万规模按需调整分块逻辑) size = 110000 dist_mat = np.full((size, size), np.inf, dtype=np.float32) neigh_idx_mat = np.full((size, size), -1, dtype=np.int32) # 假设sub_idx、sub_dist是直接生成的子块索引、子块距离数组,不需要从元组提取 rows, cols = np.ix_(index, index) mask = sub_dist < dist_mat[rows, cols] dist_mat[rows[mask], cols[mask]] = sub_dist[mask] neigh_idx_mat[rows[mask], cols[mask]] = sub_idx[mask]
内容的提问来源于stack exchange,提问作者josseossa
相关产品推荐
相关产品推荐

