如何在Python中从稀疏二维矩阵生成预定义大小的小批量数据
解决稀疏钢琴卷帘矩阵的小批量生成问题
错误原因分析
你之前的代码核心问题是混淆了非零元素索引和样本行索引:你试图按非零元素的数量来切分batch,但每个样本的非零元素数量不一致,导致切出来的元素属于原矩阵的不同行;同时新建batch矩阵时,直接使用原矩阵的行号(远大于batch的行维度),自然触发"row index exceeds matrix dimensions"错误。
而你用np.where遍历行号的方法慢,是因为它会扫描整个非零元素的行号数组,当数据量很大时,时间开销极高。
高效解决方案
方案1:直接使用CSR矩阵行切片(最推荐)
CSR矩阵原生支持高效的行切片操作,利用indptr数组可以直接定位指定行的非零元素范围,无需遍历所有数据,速度极快。
import scipy.sparse as sp # 加载磁盘上的稀疏矩阵 sparse_matrix = sp.load_npz(file_name) total_samples = sparse_matrix.shape[0] batch_size = num_samples # 你的预定义批量大小 for batch_idx in range(total_samples // batch_size): start_row = batch_idx * batch_size end_row = min((batch_idx + 1) * batch_size, total_samples) # 直接切片获取包含batch_size个样本的子稀疏矩阵 batch_sparse = sparse_matrix[start_row:end_row, :] # 按需转换格式(比如转COO)或直接用于模型输入 batch_coo = batch_sparse.tocoo() # 此时batch_coo的row是0到(batch_size-1)的本地索引,不会越界
方案2:基于COO格式的快速分块(需用COO时)
如果必须用COO格式处理,可以利用COO矩阵的row数组是按行号递增排序的特性,用np.searchsorted进行二分查找,快速定位每个batch对应的非零元素范围,替代慢遍历的np.where。
import scipy.sparse as sp import numpy as np sparse_matrix = sp.load_npz(file_name) coo_matrix = sparse_matrix.tocoo() total_samples = sparse_matrix.shape[0] batch_size = num_samples row_array = coo_matrix.row # 已排序的行号数组 for batch_idx in range(total_samples // batch_size): start_row = batch_idx * batch_size end_row = min((batch_idx + 1) * batch_size, total_samples) # 二分查找定位batch对应的非零元素起止索引 start_idx = np.searchsorted(row_array, start_row, side='left') end_idx = np.searchsorted(row_array, end_row, side='left') # 提取batch数据并转换为本地行号 batch_data = coo_matrix.data[start_idx:end_idx] batch_row = coo_matrix.row[start_idx:end_idx] - start_row # 转换为batch内的相对行号 batch_col = coo_matrix.col[start_idx:end_idx] # 构建合法的batch稀疏矩阵 batch_sparse = sp.coo_matrix( (batch_data, (batch_row, batch_col)), shape=(end_row - start_row, sparse_matrix.shape[1]) )
关键注意点
- 永远按样本行数量切分batch,而非非零元素数量,确保每个batch包含完整的样本。
- CSR行切片是最高效的方式,避免不必要的格式转换。
- 利用COO的row数组排序特性,用二分查找替代全量遍历,大幅提升速度。
内容的提问来源于stack exchange,提问作者Mateusz Dorobek
相关产品推荐
相关产品推荐

