如何在Scipy稀疏矩阵中获取两矩阵共有的非零元素并绘制散点图?
解决大规模稀疏矩阵共同非零元素散点图的高效方案
你说得太对了——当矩阵规模上去之后,转密集格式简直是灾难,内存直接爆掉不说,计算速度也慢得离谱。其实我们完全可以利用稀疏矩阵的特性,只挑出两个矩阵都不为零的元素来绘图,根本不需要碰那些占绝大多数的零元素。
下面给你两种高效的实现方式,都能完美避开密集格式的坑:
方法一:利用稀疏矩阵的布尔乘法(简洁易读)
这种方法借助scipy稀疏矩阵的内置操作,代码简洁,可读性拉满,适合大多数场景:
from scipy import sparse as sp import matplotlib.pyplot as plt import numpy as np # 生成大规模示例稀疏矩阵(10k*10k,1%非零密度) t1 = sp.random(10000, 10000, 0.01) t2 = sp.random(10000, 10000, 0.01) # 1. 标记两个矩阵共同的非零位置 # 先转布尔型稀疏矩阵,再做元素乘法——结果的非零位置就是两者都非零的位置 common_nonzero_mask = t1.astype(bool).multiply(t2.astype(bool)) # 2. 提取对应位置的数值 t1_common_vals = t1[common_nonzero_mask.nonzero()] t2_common_vals = t2[common_nonzero_mask.nonzero()] # 3. 绘制散点图 plt.scatter(t1_common_vals, t2_common_vals, marker='o', color='k', alpha=0.1) plt.xlabel('t1 Non-Zero Values') plt.ylabel('t2 Non-Zero Values') plt.title('Scatter Plot of Common Non-Zero Elements') plt.show()
方法二:位置编码找交集(极致性能)
如果你的矩阵非零元素数量特别大(比如千万级),可以用这种方法,通过位置编码减少内存开销,进一步提升速度:
from scipy import sparse as sp import matplotlib.pyplot as plt import numpy as np t1 = sp.random(10000, 10000, 0.01) t2 = sp.random(10000, 10000, 0.01) # 1. 转COO格式,方便直接获取非零元素的行、列和值 t1_coo = t1.tocoo() t2_coo = t2.tocoo() # 2. 将每个非零位置编码为唯一整数(行号*总列数 + 列号) n_cols = t1.shape[1] t1_positions = t1_coo.row * n_cols + t1_coo.col t2_positions = t2_coo.row * n_cols + t2_coo.col # 3. 找到两个位置数组的交集 common_positions = np.intersect1d(t1_positions, t2_positions, assume_unique=True) # 4. 筛选出共同位置对应的数值 t1_mask = np.isin(t1_positions, common_positions) t2_mask = np.isin(t2_positions, common_positions) t1_common_vals = t1_coo.data[t1_mask] t2_common_vals = t2_coo.data[t2_mask] # 5. 绘图 plt.scatter(t1_common_vals, t2_common_vals, marker='o', color='k', alpha=0.1) plt.xlabel('t1 Non-Zero Values') plt.ylabel('t2 Non-Zero Values') plt.title('Scatter Plot of Common Non-Zero Elements') plt.show()
为什么这两种方法更优?
两种方法都只处理非零元素,内存占用和计算量完全和非零元素的数量成正比,而不是矩阵的总大小。比如10k*10k的矩阵,总元素是1亿,但1%密度下非零元素只有100万,处理量直接降到原来的1%,效率提升非常明显。
内容的提问来源于stack exchange,提问作者mgalardini
相关产品推荐
相关产品推荐

