Python能否无需加载文件至RAM直接写入?大数据集分片洗牌咨询
嘿,这个大数据集洗牌的问题我太熟了!你遇到的痛点完全合理——np.save/savez确实搞不定这种超大规模的流式处理,而h5py刚好是为这种场景设计的。我给你一步步拆解解决方案:
核心思路
因为数据集塞不进内存,我们必须用流式处理:边按时间顺序读取源数据,边随机分配到不同分组,同时用支持增量写入的格式(HDF5)来存储分组数据,全程不用把整个数据集加载到RAM里。等所有数据分配完成后,再对每个分组单独做洗牌(同样用分块方式,避免内存溢出)。
具体实现步骤
1. 初始化分组的HDF5存储
首先我们为每个分组创建一个可扩展的HDF5数据集——HDF5允许我们动态扩大数据集的大小,不用提前知道最终数据量。
import h5py import numpy as np # 定义分组数量,比如分成3组 num_groups = 3 group_handles = [] # 存储每个分组的文件句柄和数据集对象 for group_idx in range(num_groups): # 创建分组文件 h5_file = h5py.File(f"group_{group_idx}.h5", "w") # 创建可扩展数据集:假设你的数据是二维(样本数×特征数),这里特征数设为100,按需修改 # maxshape=(None, 100)表示第一个维度(样本数)可以无限扩展 dataset = h5_file.create_dataset( "data", shape=(0, 100), maxshape=(None, 100), dtype=np.float32, # 要和源数据类型一致,按需修改 chunks=True # 开启分块存储,提升读写效率 ) group_handles.append((h5_file, dataset))
2. 流式读取源数据并分配分组
不管你的源数据是单个大HDF5文件,还是多个按时间排序的.npy文件,都可以用流式方式读取:
情况1:源数据是单个大HDF5文件
# 打开源文件(只读模式) with h5py.File("large_source_data.h5", "r") as source_file: source_dset = source_file["data"] chunk_size = 1000 # 每次读取的样本数,根据你的内存大小调整 # 按块遍历源数据 for start_idx in range(0, source_dset.shape[0], chunk_size): end_idx = min(start_idx + chunk_size, source_dset.shape[0]) # 读取当前块的数据(这部分只会占用chunk_size大小的内存) current_chunk = source_dset[start_idx:end_idx] # 为当前块的每个样本随机分配分组 group_assignments = np.random.randint(0, num_groups, size=current_chunk.shape[0]) # 将数据写入对应分组 for g_idx in range(num_groups): # 筛选属于当前分组的样本 group_chunk = current_chunk[group_assignments == g_idx] if len(group_chunk) == 0: continue # 获取分组的数据集对象 _, target_dset = group_handles[g_idx] # 扩展数据集的大小 current_size = target_dset.shape[0] target_dset.resize(current_size + len(group_chunk), axis=0) # 写入数据 target_dset[current_size:] = group_chunk
情况2:源数据是多个按时间排序的.npy文件
import glob # 按时间顺序获取所有.npy文件(假设文件名带时间戳,比如data_20230101.npy) npy_files = sorted(glob.glob("data_*.npy")) for file_path in npy_files: # 用内存映射读取.npy文件,不用加载全部到内存 with np.load(file_path, mmap_mode="r") as source_arr: chunk_size = 1000 for start_idx in range(0, source_arr.shape[0], chunk_size): end_idx = min(start_idx + chunk_size, source_arr.shape[0]) current_chunk = source_arr[start_idx:end_idx] group_assignments = np.random.randint(0, num_groups, size=current_chunk.shape[0]) for g_idx in range(num_groups): group_chunk = current_chunk[group_assignments == g_idx] if len(group_chunk) == 0: continue _, target_dset = group_handles[g_idx] current_size = target_dset.shape[0] target_dset.resize(current_size + len(group_chunk), axis=0) target_dset[current_size:] = group_chunk
3. 关闭分组文件
所有数据分配完成后,记得关闭文件句柄:
for h5_file, _ in group_handles: h5_file.close()
分组后洗牌:分块处理避免内存溢出
每个分组的HDF5文件还是很大,不能直接加载到内存洗牌,我们可以用分块洗牌的方法,或者原地交换块来实现:
方法1:临时数据集法(需要额外磁盘空间,更简单)
def shuffle_hdf5_group(file_path): with h5py.File(file_path, "r+") as h5_file: dset = h5_file["data"] total_samples = dset.shape[0] # 生成全局随机排列的索引 permutation = np.random.permutation(total_samples) chunk_size = 1000 # 创建临时存储数据集 temp_dset = h5_file.create_dataset( "temp_data", shape=dset.shape, dtype=dset.dtype, chunks=True ) # 分块写入打乱后的数据 for start_idx in range(0, total_samples, chunk_size): end_idx = min(start_idx + chunk_size, total_samples) temp_dset[start_idx:end_idx] = dset[permutation[start_idx:end_idx]] # 替换原数据集 del h5_file["data"] h5_file.move("temp_data", "data")
方法2:原地块交换法(无需额外磁盘空间,效率稍低)
def shuffle_hdf5_group_inplace(file_path): with h5py.File(file_path, "r+") as h5_file: dset = h5_file["data"] total_samples = dset.shape[0] chunk_size = 1000 num_chunks = (total_samples + chunk_size - 1) // chunk_size # 生成块的随机排列 chunk_perm = np.random.permutation(num_chunks) # 创建缓冲区存储交换的块 buffer = np.zeros((chunk_size, dset.shape[1]), dtype=dset.dtype) for i in range(num_chunks): if chunk_perm[i] == i: continue # 读取当前块到缓冲区 current_chunk_len = min(chunk_size, total_samples - i*chunk_size) buffer[:current_chunk_len] = dset[i*chunk_size : i*chunk_size + current_chunk_len] # 把目标块写入当前位置 target_chunk_len = min(chunk_size, total_samples - chunk_perm[i]*chunk_size) dset[i*chunk_size : i*chunk_size + target_chunk_len] = dset[chunk_perm[i]*chunk_size : chunk_perm[i]*chunk_size + target_chunk_len] # 把缓冲区内容写入目标位置 dset[chunk_perm[i]*chunk_size : chunk_perm[i]*chunk_size + current_chunk_len] = buffer[:current_chunk_len]
为什么不用numpy的save/savez?
np.save只能一次性写入整个数组,无法动态扩展,你没法边读边写;savez是把多个数组打包成压缩文件,同样需要把所有数据加载到内存才能生成;- 两者都不支持内存映射式的增量读写,完全不适合超大数据集场景。
而h5py的优势就在于:支持分块存储、动态扩展数据集、内存映射读取,完美匹配你的需求。
内容的提问来源于stack exchange,提问作者Charlie
相关产品推荐
相关产品推荐

