Python分批读取与处理大型稀疏矩阵的方案问询
大型稀疏矩阵分批处理方案(无需全量加载)
问题场景
存储在文件中的大型稀疏矩阵,受内存限制无法全量加载,需要实现高效的分批读取与处理,现有基于scipy.sparse的方案会全量加载矩阵,无法满足需求。
可行解决方案
1. 改用HDF5格式实现分块读取
npz格式不支持部分读取,因此先将稀疏矩阵转存为HDF5格式,通过分块存储实现按需加载指定行片段:
预处理:将稀疏矩阵保存为HDF5
import h5py import scipy.sparse as sp import numpy as np def save_sparse_to_hdf5(matrix, save_path, chunk_rows=1000): with h5py.File(save_path, 'w') as f: # 存储矩阵元信息 f.attrs['shape'] = matrix.shape f.attrs['format'] = matrix.format # 按行分块存储为COO格式 for chunk_idx in range(0, matrix.shape[0], chunk_rows): end_idx = min(chunk_idx + chunk_rows, matrix.shape[0]) chunk = matrix[chunk_idx:end_idx].tocoo() # 创建chunk分组,存储行、列、数据 grp = f.create_group(f'chunk_{chunk_idx//chunk_rows}') grp.create_dataset('row', data=chunk.row + chunk_idx) # 全局行索引 grp.create_dataset('col', data=chunk.col) grp.create_dataset('data', data=chunk.data)
分批加载的DataLoader实现
class SparseHDF5Loader: def __init__(self, hdf5_path, batch_size, image_width, image_height): self.hdf5_file = h5py.File(hdf5_path, 'r') self.total_rows = self.hdf5_file.attrs['shape'][0] self.total_cols = self.hdf5_file.attrs['shape'][1] self.batch_size = batch_size self.image_width = image_width self.image_height = image_height self.rows_per_batch = batch_size * image_width self.num_batches = self.total_rows // self.rows_per_batch # 从第一个chunk名获取分块大小,也可提前存入attrs self.chunk_rows = int(list(self.hdf5_file.keys())[0].split('_')[1]) * 1000 def __iter__(self): for batch_idx in range(self.num_batches): start_row = batch_idx * self.rows_per_batch end_row = start_row + self.rows_per_batch # 确定需要加载的chunk范围 start_chunk_idx = start_row // self.chunk_rows end_chunk_idx = (end_row - 1) // self.chunk_rows # 收集当前批次的稀疏数据 rows, cols, data = [], [], [] for chunk_idx in range(start_chunk_idx, end_chunk_idx + 1): chunk_grp = self.hdf5_file[f'chunk_{chunk_idx}'] # 筛选属于当前批次的行 mask = (chunk_grp['row'][...] >= start_row) & (chunk_grp['row'][...] < end_row) # 转为批次内的本地行索引 rows.append(chunk_grp['row'][...][mask] - start_row) cols.append(chunk_grp['col'][...][mask]) data.append(chunk_grp['data'][...][mask]) # 构建批次稀疏矩阵并转为图像格式 batch_sparse = sp.coo_matrix( (np.concatenate(data), (np.concatenate(rows), np.concatenate(cols))), shape=(self.rows_per_batch, self.total_cols) ) batch_dense = batch_sparse.toarray() batch_images = batch_dense.reshape(self.batch_size, self.image_width, self.image_height, 1) batch_images = np.transpose(batch_images, (0, 2, 1, 3)) yield batch_images def __del__(self): self.hdf5_file.close()
2. 使用Dask Sparse实现延迟加载
Dask支持稀疏矩阵的分块延迟计算,无需全量加载,直接基于原npz文件实现分批处理:
import dask.sparse as ds import numpy as np class DaskSparseLoader: def __init__(self, npz_path, batch_size, image_width, image_height): self.sparse_matrix = ds.load_npz(npz_path) self.total_rows = self.sparse_matrix.shape[0] self.batch_size = batch_size self.image_width = image_width self.image_height = image_height self.rows_per_batch = batch_size * image_width self.num_batches = self.total_rows // self.rows_per_batch def __iter__(self): for batch_idx in range(self.num_batches): start_row = batch_idx * self.rows_per_batch end_row = start_row + self.rows_per_batch # 切片获取批次(Dask延迟加载,仅在compute时读取对应数据) batch_sparse = self.sparse_matrix[start_row:end_row, :] batch_dense = batch_sparse.compute() # 转换为图像格式 batch_images = batch_dense.reshape(self.batch_size, self.image_width, self.image_height, 1) batch_images = np.transpose(batch_images, (0, 2, 1, 3)) yield batch_images
3. 预分块存储为多个npz文件
如果不想修改存储格式,可提前将大矩阵按行分割为多个小npz文件,读取时按需加载对应文件:
预处理:分割稀疏矩阵为多个npz文件
import scipy.sparse as sp def split_sparse_to_npz_chunks(matrix, base_save_path, rows_per_chunk=1000): # 保存总行数到配置文件 with open(f'{base_save_path}_shape.txt', 'w') as f: f.write(f"{matrix.shape[0]}\n{matrix.shape[1]}") # 分块保存 for chunk_idx in range(0, matrix.shape[0], rows_per_chunk): end_idx = min(chunk_idx + rows_per_chunk, matrix.shape[0]) chunk = matrix[chunk_idx:end_idx] sp.save_npz(f'{base_save_path}_chunk_{chunk_idx//rows_per_chunk}.npz', chunk)
分批加载的DataLoader实现
import scipy.sparse as sp import numpy as np class SplitSparseLoader: def __init__(self, base_path, batch_size, image_width, image_height, rows_per_chunk=1000): # 读取总行列数 with open(f'{base_path}_shape.txt', 'r') as f: self.total_rows = int(f.readline().strip()) self.total_cols = int(f.readline().strip()) self.base_path = base_path self.batch_size = batch_size self.image_width = image_width self.image_height = image_height self.rows_per_batch = batch_size * image_width self.rows_per_chunk = rows_per_chunk self.num_batches = self.total_rows // self.rows_per_batch def __iter__(self): for batch_idx in range(self.num_batches): start_row = batch_idx * self.rows_per_batch end_row = start_row + self.rows_per_batch # 确定需要加载的chunk文件范围 start_chunk_idx = start_row // self.rows_per_chunk end_chunk_idx = (end_row - 1) // self.rows_per_chunk # 加载并合并对应chunk的目标行 batch_chunks = [] for chunk_idx in range(start_chunk_idx, end_chunk_idx + 1): chunk = sp.load_npz(f'{self.base_path}_chunk_{chunk_idx}.npz') # 计算当前chunk中需要截取的本地行范围 local_start = max(start_row - chunk_idx * self.rows_per_chunk, 0) local_end = min(end_row - chunk_idx * self.rows_per_chunk, self.rows_per_chunk) batch_chunks.append(chunk[local_start:local_end]) # 合并为批次稀疏矩阵并转换格式 batch_sparse = sp.vstack(batch_chunks) batch_dense = batch_sparse.toarray() batch_images = batch_dense.reshape(self.batch_size, self.image_width, self.image_height, 1) batch_images = np.transpose(batch_images, (0, 2, 1, 3)) yield batch_images
内容的提问来源于stack exchange,提问作者Mateusz Dorobek
相关产品推荐
相关产品推荐

