如何修改Scipy稀疏矩阵的指定子矩阵?并实现循环向量化加速
问题根源
Scipy的CSR矩阵切片返回的是副本而非视图,你原来的A[tk][:, tk] += ...操作只会修改临时副本,根本不会作用到原矩阵上,这就是A始终为零矩阵的原因。
解决方案
方案1:循环内正确修改CSR矩阵
CSR矩阵不支持直接切片赋值,必须明确指定要修改的行/列索引和对应的增量值:
import scipy.sparse as sp A = sp.csr_matrix((n, n)) for k in range(nt): tk = t[k] xk = x[tk] a = function1(xk) # 生成3×3子矩阵对应的所有行、列索引对 rows = tk.repeat(3) # 结果:[tk0, tk0, tk0, tk1, tk1, tk1, tk2, tk2, tk2] cols = np.tile(tk, 3)# 结果:[tk0, tk1, tk2, tk0, tk1, tk2, tk0, tk1, tk2] # 获取增量值并扁平化 delta = a.dot(a.T).ravel() # 直接更新原矩阵的对应位置 A[rows, cols] += delta
如果tk中存在重复索引,该方式会自动累加对应位置的值,符合需求。
方案2:向量化批量更新(推荐,大幅提升速度)
Python循环效率极低,尤其当nt很大时,推荐用向量化方式一次性构建所有增量再合并到A中:
import scipy.sparse as sp import numpy as np # 1. 生成所有子矩阵的行、列索引对 rows = t.repeat(3, axis=0).ravel() # 每个tk的3个元素各重复3次后扁平化 cols = np.tile(t, (3,1)).ravel() # 每个tk重复3行后扁平化 # 2. 批量计算所有增量值 xk_batch = x[t] # 获取所有tk对应的xk,得到(m,3)数组 # 批量计算每个xk对应的a矩阵,得到(m,3,3)数组(若function1支持批量输入可直接调用) a_batch = np.array([function1(xk) for xk in xk_batch]) # 扁平化每个a.dot(a.T),得到(m*9,)的增量数组 delta_batch = (a_batch @ a_batch.transpose(0,2,1)).ravel() # 3. 构建增量稀疏矩阵并合并到A delta_coo = sp.coo_matrix((delta_batch, (rows, cols)), shape=(n, n)) A = A + delta_coo
关键提示
- 若
function1支持批量输入(直接接收(m,3)数组返回(m,3,3)矩阵),可去掉列表推导式直接调用,速度会进一步提升。 - COO格式是构建稀疏矩阵增量最高效的格式,仅需一次结构更新,远快于循环内修改CSR矩阵。
内容的提问来源于stack exchange,提问作者qwe123
相关产品推荐
相关产品推荐

