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

如何基于含多批次的多文件创建PyTorch Dataset?

处理多文件批量数据集的PyTorch工具方案

1. PyTorch官方:torch.utils.data.IterableDataset

这是官方原生适配流式/多文件批量数据的方案,无需一次性加载全量数据。

  • 核心逻辑:在__iter__方法中逐个遍历文件,读取单文件内的批量数据后拆分返回单样本
  • Parquet场景示例代码:
import torch
import pandas as pd
from torch.utils.data import IterableDataset, DataLoader

class ParquetIterableDataset(IterableDataset):
    def __init__(self, file_paths):
        self.file_paths = file_paths

    def __iter__(self):
        for file_path in self.file_paths:
            df = pd.read_parquet(file_path)
            # 按需求返回文本或嵌入向量样本
            for _, row in df.iterrows():
                yield torch.tensor(row['embedding']), row['text']

# 使用示例
file_list = ["data/part_0.parquet", "data/part_1.parquet"]
dataset = ParquetIterableDataset(file_list)
dataloader = DataLoader(dataset, batch_size=32)
  • 优势:官方原生支持,自定义性强,可灵活适配不同文件格式的读取逻辑
  • 注意:多进程加载时需在__iter__中处理文件分配,避免进程间重复读取

2. 第三方成熟库:Hugging Face datasets

NLP领域主流数据集处理库,原生支持Parquet分区、多文件批量数据,无缝对接PyTorch DataLoader。

  • 核心特性:自动识别分区结构,支持流式加载(低内存占用),内置预处理管道
  • 示例代码:
from datasets import load_dataset

# 加载Parquet分区数据集,支持通配符匹配文件
dataset = load_dataset("parquet", data_files={"train": "data/part_*.parquet"}, streaming=True)

# 转换为PyTorch兼容格式
torch_dataset = dataset['train'].with_format("torch")
dataloader = torch.utils.data.DataLoader(torch_dataset, batch_size=32)
  • 优势:无需手动实现文件遍历,支持多种数据格式,内置缓存、数据增强等功能,适配大规模文本/嵌入数据集
  • 注意:streaming模式为流式读取,适合超大数据集;非流式模式会加载全量数据到内存

3. 第三方库:PyTorch Lightning生态工具

若使用PyTorch Lightning训练,其lightning.data模块(基于原torchdata重构)提供了便捷的文件遍历与批量读取工具,兼容PyTorch Dataset规范。

  • 核心逻辑:通过from_file_system遍历文件,结合map操作读取单文件批量数据
  • 示例代码片段:
from lightning.data import from_file_system

# 遍历指定目录下的所有Parquet文件
file_dataset = from_file_system("data/", filters=["*.parquet"])
# 读取每个文件并拆分样本
dataset = file_dataset.map(lambda path: pd.read_parquet(path).iterrows())
  • 优势:继承了torchdata的便捷性,维护状态活跃,适合Lightning体系用户

内容的提问来源于stack exchange,提问作者dule arnaux

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 11:43:14