HuggingFace:借助自定义data_loader与data_collator实现本地数据集流式加载并解决跨系统缓存问题
解决Hugging Face流式加载数据集避免缓存的问题
针对你遇到的跨系统缓存无法复用、想要用流式加载替代全量缓存的需求,我来一步步帮你梳理解决方案,同时解答你关于__len__和__iter__的疑问:
核心原理
流式加载(streaming=True)的本质是按需生成/加载样本,不会将全量数据缓存为.arrow文件,而是在训练时实时生成并处理样本,完美解决跨系统缓存不兼容的问题。你的自定义GeneratorBasedBuilder其实已经天然适配流式模式,只需要做一些小调整即可。
1. 自定义数据集Builder的适配(无需手动写__iter__)
你的GeneratorBasedBuilder类中的_generate_examples本身就是一个生成器函数(通过yield返回样本),当开启streaming=True时,Hugging Face Datasets库会自动将其包装成可迭代的流式Dataset对象,你完全不需要手动实现__iter__方法。
不过要确保_generate_examples是标准的生成器逻辑:
# my_data_loader.py from datasets import GeneratorBasedBuilder, DatasetInfo, SplitGenerator, Split class MyCustomDataset(GeneratorBasedBuilder): def _info(self): # 定义你的数据集特征,比如文本、标签等 return DatasetInfo( features=... # 替换成你的特征定义,例如 Features({"text": Value("string")}) ) def _split_generators(self, dl_manager): # 返回训练/验证集的生成配置,无需处理缓存 return [ SplitGenerator( name=Split.TRAIN, gen_kwargs={"data_source": "/path/to/your/train_data"} # 传递数据来源参数 ), SplitGenerator( name=Split.VALIDATION, gen_kwargs={"data_source": "/path/to/your/val_data"} ) ] def _generate_examples(self, data_source): # 核心生成器:逐个yield样本(这里以读取文本文件为例) with open(data_source, "r", encoding="utf-8") as f: for sample_id, line in enumerate(f): # 处理单条数据成样本字典 processed_sample = {"text": line.strip()} yield sample_id, processed_sample
2. 流式数据集的处理与训练代码调整
在主训练脚本中,你需要确保map、shuffle等操作都是惰性执行的(流式模式下默认就是惰性的),同时适配Trainer的配置:
from datasets import load_dataset from transformers import Trainer, TrainingArguments from your_module import MyDataCollator, model, tokenizer # 替换成你的实际模块 # 流式加载自定义数据集 dataset = load_dataset("./my_data_loader.py", streaming=True) train_dataset = dataset["train"] # 惰性执行分词映射:流式下支持批量处理提高效率 train_dataset = train_dataset.map( lambda batch: tokenizer( batch["text"], truncation=True, padding="max_length", max_length=512 ), batched=True, # 批量处理,建议设置合理的batch_size batch_size=1000 ) # 可选:流式洗牌(需要设置buffer_size,越大随机性越好,内存占用越高) train_dataset = train_dataset.shuffle(buffer_size=10000) # 初始化自定义DataCollator(确保它能处理批量样本) data_collator = MyDataCollator(...) # 配置训练参数:注意流式下的特殊设置 training_args = TrainingArguments( output_dir="./training_results", per_device_train_batch_size=8, max_steps=10000, # 流式下推荐用max_steps代替num_train_epochs(如果不知道总样本数) logging_dir="./logs", logging_steps=100, save_steps=500 ) # 初始化Trainer:直接传入流式数据集即可 trainer = Trainer( model=model, args=training_args, train_dataset=train_dataset, data_collator=data_collator, # 如果需要验证,也要传入流式验证集 # eval_dataset=dataset["validation"].map(...).shuffle(buffer_size=5000) ) # 启动训练:全程无全量缓存,实时生成样本 trainer.train()
关于__len__的疑问解答
流式数据集默认没有__len__方法,因为数据是实时生成的,无法提前获取总样本数。但你有两种方式适配Trainer的步数计算:
- 手动指定样本数:在
_split_generators的SplitGenerator中添加num_examples参数,这样即使是流式加载,Dataset也会返回这个预设的样本数,Trainer可以据此计算总epoch步数:
def _split_generators(self, dl_manager): return [ SplitGenerator( name=Split.TRAIN, gen_kwargs={"data_source": "/path/to/train_data"}, num_examples=100000 # 手动设置训练集总样本数 ) ]
- 用
max_steps控制训练时长:如果不知道总样本数,直接在TrainingArguments中设置max_steps,训练到指定步数就停止,无需依赖数据集长度。
流式加载的注意事项
- 避免使用需要全量数据集的操作:比如
dataset["train"].sort()、dataset["train"].unique()这类需要遍历全量数据的方法,流式下不支持。 - 批量处理优化:
map时设置batched=True能大幅提高分词等预处理的效率,建议根据内存情况调整batch_size。 - 跨系统兼容性:流式加载不需要生成任何本地缓存文件,只要其他系统能访问到你的数据来源(比如相同的文件路径、数据库连接),就能直接运行训练代码,无需重新缓存。
内容的提问来源于stack exchange,提问作者Mohbat Tharani
相关产品推荐
相关产品推荐

