如何基于含多批次的多文件创建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
相关产品推荐
相关产品推荐

