You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何修改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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.28 03:57:24