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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 11:45:44