如何在Hugging Face数据集流式模式下使用.map()并保持迭代器特性
关于Hugging Face Datasets流式处理与数据集交错的问题
我正在使用Hugging Face datasets库,需要对多个数据集(如ds_khan、ds_mathematica)应用.map()函数,同时以流式方式(不加载全量数据到内存)处理,还要交错转换后的数据集并尽可能延迟计算,达到类似streaming=True的效果。
当前代码如下:
from datasets import load_dataset, interleave_datasets import os # 补充原代码缺失的导入 def get_hf_khan_ds(path_2_ds: str, split: str = 'train'): path_2_ds = os.path.expanduser(path_2_ds) dataset = load_dataset('json', data_files=[path_2_ds], split=split, streaming=True) problem_as_text = lambda example: {'text': example['problem']} return dataset.map(problem_as_text, remove_columns=dataset.column_names) def main(): ds_khan = get_hf_khan_ds('~/gold-ai-olympiad/data/amps/khan/train.jsonl') ds_mathematica = get_hf_khan_ds('~/gold-ai-olympiad/data/amps/mathematica/train.jsonl') interleaved_datasets = interleave_datasets([ds_khan, ds_mathematica], probabilities=[0.5, 0.5]) for sample in interleaved_datasets.take(10): print(sample) if __name__ == '__main__': main()
但运行时出现全量数据处理进度条,加载耗时远超预期,不确定是否正确实现流式与延迟计算,现提出问题:
- 该代码是否正确实现流式/迭代器式转换?
- 若否,如何修改以确保仅按需处理数据,不预加载全量内容?
- 在流式模式下,有无更高效的数据集交错方式?
注:数据存储在本地,最终将使用HF数据集。
问题解答
1. 该代码是否正确实现流式/迭代器式转换?
没有正确实现。默认情况下,map()在流式数据集上会触发全量预处理(出现进度条就是明确信号),它会提前遍历整个数据集完成转换,没有做到真正的延迟计算,违背了流式处理的初衷。
2. 如何修改以确保仅按需处理数据?
核心是让map()以逐样本的延迟方式执行,修改要点:
- 在
map()调用中添加batched=False参数(默认batched=True会触发批量全量扫描) - 避免任何会触发全量数据扫描的操作(比如调用
len()、统计数据集信息等)
修改后的代码:
from datasets import load_dataset, interleave_datasets import os def get_hf_khan_ds(path_2_ds: str, split: str = 'train'): path_2_ds = os.path.expanduser(path_2_ds) dataset = load_dataset('json', data_files=[path_2_ds], split=split, streaming=True) problem_as_text = lambda example: {'text': example['problem']} # 添加batched=False,确保逐样本延迟处理 return dataset.map(problem_as_text, remove_columns=dataset.column_names, batched=False) def main(): ds_khan = get_hf_khan_ds('~/gold-ai-olympiad/data/amps/khan/train.jsonl') ds_mathematica = get_hf_khan_ds('~/gold-ai-olympiad/data/amps/mathematica/train.jsonl') interleaved_datasets = interleave_datasets([ds_khan, ds_mathematica], probabilities=[0.5, 0.5]) for sample in interleaved_datasets.take(10): print(sample) if __name__ == '__main__': main()
这样修改后,map()会在迭代到对应样本时才执行转换,不会提前全量处理数据。
3. 流式模式下更高效的数据集交错方式
- 优先使用官方
interleave_datasets:只要输入的数据集是纯流式迭代器(已按上述修改实现延迟计算),这个方法是最简洁高效的,还能通过probabilities控制各数据集的采样权重,seed参数保证交错结果可复现。 - 手动实现迭代器交错:如果需要固定轮次交替(如先取一个ds_khan样本,再取一个ds_mathematica样本),可以手动写迭代器逻辑,避免权重计算的额外开销,示例:
def interleave_streaming_datasets(datasets): iterators = [iter(ds) for ds in datasets] while True: for it in iterators: try: yield next(it) except StopIteration: pass # 在main函数中替换interleave_datasets调用 interleaved_datasets = interleave_streaming_datasets([ds_khan, ds_mathematica])
这种方式适合对交错顺序有严格要求的场景,性能更优。
内容的提问来源于stack exchange,提问作者Charlie Parker
相关产品推荐
相关产品推荐

