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

PyTorch自定义Dataset加载HDF5数据时DataLoader速度过慢问题

解决PyTorch DataLoader加载HDF5数据集速度过慢的问题

问题根源分析

你的DataLoader加载速度比直接用h5py慢100倍,核心原因有三个:

  • 单样本随机读取的高开销:__getitem__每次只读取单个样本,而HDF5是块存储格式,随机读取小数据的IO开销远大于批量读取。你直接测试用的是批量切片dataset[a:100*a],天然利用了批量读取的高效性。
  • 多进程下HDF5文件句柄冲突:DataLoader默认启用多进程,但h5py的File对象无法安全地在多进程间共享。子进程继承父进程的文件句柄后,会出现IO阻塞或重复打开文件的情况,大幅降低效率。
  • 单样本转Tensor的累积开销:每次单独将单个样本转为Tensor,相比批量转换的开销会不断累积,进一步拖慢速度。

针对性解决方案

方案1:修复多进程下的HDF5文件访问问题

修改Dataset,让每个worker进程单独打开自己的HDF5文件句柄,避免跨进程共享的冲突,同时开启persistent_workers复用worker进程:

import time
from utils import get_vocab_and_skipgrams
from torch.utils.data import Dataset
from torch.utils.data import DataLoader
import os
import h5py
import numpy as np
import torch

class CustomSkipGramDataset(Dataset):
    def __init__(self, filename, window_size, data_dir="training_data", data_exists=True):
        self.window_size = window_size
        self.filename = filename
        self.data_exists = data_exists
        self.vocab_path = os.path.join(data_dir, "vocab.npy")
        self.hdf5_path = os.path.join(data_dir, "skipgram.h5")
        
        if not data_exists:
            get_vocab_and_skipgrams(filename, data_dir)
        
        self.vocab = np.load(self.vocab_path, allow_pickle=True).tolist()
        self.vocab_size = len(self.vocab)
        # 不在初始化阶段打开文件,留给worker进程处理
        self.hf = None
        self.dataset = None
        
        # 提前获取数据集长度,避免每次__len__都打开文件
        with h5py.File(self.hdf5_path, "r") as temp_hf:
            self.total_samples = temp_hf["positive_skips"].shape[0]
    
    def __len__(self):
        return self.total_samples
    
    def __getitem__(self, index):
        # 每个worker进程首次调用时打开自己的文件句柄
        if self.hf is None:
            self.hf = h5py.File(self.hdf5_path, "r")
            self.dataset = self.hf["positive_skips"]
        
        x, y = self.dataset[index]
        # 用torch.from_numpy替代torch.tensor,共享内存提升速度
        return torch.from_numpy(x).long(), torch.from_numpy(y).long()

使用DataLoader时配置参数:

dataset = CustomSkipGramDataset("your_filename", window_size=2)
# 开启多进程+持久化worker,避免重复打开文件
dataloader = DataLoader(
    dataset,
    batch_size=32,
    num_workers=4,  # 根据CPU核心数调整
    persistent_workers=True,
    pin_memory=True  # 如果用GPU,开启这个加速数据传输到CUDA
)

方案2:改用IterableDataset实现批量读取

对于无法全量加载的大型数据集,IterableDataset更适合流式批量读取,直接利用h5py的批量读取优势:

class CustomSkipGramIterableDataset(torch.utils.data.IterableDataset):
    def __init__(self, filename, window_size, batch_size, data_dir="training_data", data_exists=True):
        self.window_size = window_size
        self.filename = filename
        self.data_exists = data_exists
        self.batch_size = batch_size
        self.vocab_path = os.path.join(data_dir, "vocab.npy")
        self.hdf5_path = os.path.join(data_dir, "skipgram.h5")
        
        if not data_exists:
            get_vocab_and_skipgrams(filename, data_dir)
        
        self.vocab = np.load(self.vocab_path, allow_pickle=True).tolist()
        self.vocab_size = len(self.vocab)
        
        # 提前获取总样本数
        with h5py.File(self.hdf5_path, "r") as hf:
            self.total_samples = hf["positive_skips"].shape[0]

    def __iter__(self):
        # 多进程下划分数据给每个worker
        worker_info = torch.utils.data.get_worker_info()
        if worker_info is None:
            start_idx, end_idx = 0, self.total_samples
        else:
            per_worker = self.total_samples // worker_info.num_workers
            start_idx = worker_info.id * per_worker
            end_idx = start_idx + per_worker
            # 最后一个worker处理剩余数据
            if worker_info.id == worker_info.num_workers - 1:
                end_idx = self.total_samples
        
        # 每个worker单独打开文件句柄
        with h5py.File(self.hdf5_path, "r") as hf:
            dataset = hf["positive_skips"]
            # 按批次读取数据
            for idx in range(start_idx, end_idx, self.batch_size):
                batch_end = min(idx + self.batch_size, end_idx)
                batch_data = dataset[idx:batch_end]
                # 批量转换为Tensor,大幅提升效率
                x = torch.from_numpy(batch_data[:, 0]).long()
                y = torch.from_numpy(batch_data[:, 1]).long()
                yield x, y

使用方式:

dataset = CustomSkipGramIterableDataset("your_filename", window_size=2, batch_size=32)
# 注意batch_size设为None,因为IterableDataset已经返回批次
dataloader = DataLoader(
    dataset,
    batch_size=None,
    num_workers=4,
    persistent_workers=True,
    pin_memory=True
)

额外优化:HDF5文件本身的存储优化

如果你的HDF5文件是自己生成的,创建数据集时可以做以下优化:

  • 设置合适的chunks大小,让数据块和你的批量读取大小匹配,比如chunks=(batch_size, 2)
  • 关闭不必要的压缩(compression=None),压缩会增加读取时的解压开销

测试验证

修改后可以对比两种方式的速度:

  • 用DataLoader遍历100个批次,统计总耗时
  • 用直接批量读取的方式做同样的操作,对比耗时差距

内容的提问来源于stack exchange,提问作者Jash Jasani

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 22:50:12