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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.13 09:28:12