如何高效将稀疏矩阵单列添加到另一稀疏矩阵的所有列上
高效实现稀疏矩阵列广播加法的方案
你需要实现的是将稀疏列向量B广播为与A同维度的矩阵(每行值等于B对应行的值)后与稀疏矩阵A相加的操作,原有代码效率低的核心问题是存在O(nnz(A)*nnz(B))的嵌套循环、以及逐列检查存在性的O(nnz(B)*M)操作(M为矩阵列数),大矩阵场景下完全不可用。
首选实现方案(调用scipy内置优化运算)
scipy稀疏矩阵运算底层为C实现,效率比纯Python实现高2~3个数量级,你要的操作可以通过广播乘法直接实现:
from scipy.sparse import coo_matrix, csr_matrix import numpy as np def sparse_add_broadcast(coo_a: coo_matrix, coo_b: coo_matrix) -> coo_matrix: # 入参校验 assert coo_b.shape[1] == 1 assert coo_a.shape[0] == coo_b.shape[0] # 列向量B乘全1行向量,得到每行值为B对应行值的同维度矩阵 broadcast_b = coo_b.dot(csr_matrix(np.ones((1, coo_a.shape[1]), dtype=coo_b.dtype))) # 直接相加,scipy自动处理重复坐标的求和逻辑 return (coo_a + broadcast_b).tocoo()
自定义三元组实现方案
如果需要完全基于三元组逻辑实现,可使用以下时间复杂度为O(nnz(A) + nnz(B) + 新增非零元数量)的方案,完全避免嵌套循环:
from scipy.sparse import coo_matrix import numpy as np from collections import defaultdict def custom_sparse_add_broadcast(coo_a: coo_matrix, coo_b: coo_matrix) -> coo_matrix: n_rows, n_cols = coo_a.shape assert coo_b.shape == (n_rows, 1), "B必须是与A行数相同的列向量" # 1. 构建行对应B值的数组,时间复杂度O(nnz(B)) bval_arr = np.zeros(n_rows, dtype=coo_a.dtype) bval_arr[coo_b.row] = coo_b.data b_nonzero_rows = coo_b.row[bval_arr[coo_b.row] != 0] # 2. 处理A原有三元组的加法,向量化实现无Python循环,时间复杂度O(nnz(A)) new_data = coo_a.data + bval_arr[coo_a.row] new_rows = coo_a.row.copy() new_cols = coo_a.col.copy() # 3. 统计A中每行已有列,时间复杂度O(nnz(A)) row_exist_cols = defaultdict(set) for r, c in zip(coo_a.row, coo_a.col): row_exist_cols[r].add(c) # 4. 补充B非零行中A未覆盖的列 all_cols = set(range(n_cols)) add_rows = [] add_cols = [] add_data = [] for r in b_nonzero_rows: exist_cols = row_exist_cols.get(r, set()) missing_cols = all_cols - exist_cols b_val = bval_arr[r] for c in missing_cols: add_rows.append(r) add_cols.append(c) add_data.append(b_val) # 合并所有三元组生成结果 final_rows = np.concatenate([new_rows, add_rows]) final_cols = np.concatenate([new_cols, add_cols]) final_data = np.concatenate([new_data, add_data]) return coo_matrix((final_data, (final_rows, final_cols)), shape=(n_rows, n_cols))
注意事项
- 如果你的矩阵列数极大,且B的非零行很多,该操作本身会生成大量非零元,属于操作的固有开销,无法通过算法优化避免
- 如果你确定不需要保留原矩阵为0的列加B后的值,可省略第四步的补全操作,仅处理A原有三元组的加法即可,性能会提升数倍
内容的提问来源于stack exchange,提问作者Hart
相关产品推荐
相关产品推荐

