如何流式加载Hugging Face大数据集到Apache Beam管道避免内存溢出?
解决Hugging Face大型数据集接入Apache Beam的内存溢出问题
核心问题在于beam.Create(dataset['train'])会一次性遍历并加载整个数据集到内存,这对Oscar2301这类TB级别的数据集完全不可行。以下是几种简便的流式加载方案,避免内存溢出:
方案1:用beam.ReadFromGenerator配合Dataset惰性迭代
Hugging Face的Dataset(包括本地缓存的非流式数据集)本身支持惰性迭代,不会一次性加载全部数据。用beam.ReadFromGenerator替代beam.Create,可以让Beam逐元素拉取数据:
pipeline_args = PipelineOptions( runner='FlinkRunner', streaming=True, ) # 可选:如果数据集还未设置流式模式,可开启(本地缓存数据集也支持惰性迭代) # dataset = load_dataset("your_dataset_name", streaming=True) with beam.Pipeline(options=pipeline_args) as pipeline: preprocessed = (pipeline | beam.ReadFromGenerator(lambda: dataset['train']) # 惰性迭代数据集 | beam.Map(prepare_oscar2301_for_c4) ) # 后续处理流程保持不变 pages = preprocessed deduped_pages = c4_utils.remove_duplicate_text(pages) encoded = deduped_pages | beam.Map(lambda page: json.dumps(page.__dict__, ensure_ascii=False).encode('utf-8')) encoded | 'Write to JSONL' >> beam.io.WriteToText('/content/output.jsonl') pipeline.run()
优化细节
- 可通过
dataset['train'].set_format("arrow", batch_size=500)调整每次读取的批次大小,平衡内存占用与处理效率。 - 惰性迭代模式下,Dataset只会从本地Arrow文件中加载当前需要的批次,不会预加载全部数据。
方案2:直接读取缓存目录的Arrow文件
如果知道数据集的缓存路径,直接读取底层Arrow文件可以进一步降低内存开销,绕过Dataset的API包装:
import pyarrow as pa import os def load_arrow_files(cache_dir): # 遍历缓存目录下所有.arrow文件 for root, _, files in os.walk(cache_dir): for file in files: if file.endswith('.arrow'): table = pa.read_table(os.path.join(root, file)) # 将Arrow表逐行转为字典(与Dataset输出格式一致) yield from table.to_pylist() # 替换为你的Oscar2301数据集缓存目录 CACHE_DIR = "/path/to/huggingface/datasets/oscar2301/train" pipeline_args = PipelineOptions( runner='FlinkRunner', streaming=True, ) with beam.Pipeline(options=pipeline_args) as pipeline: preprocessed = (pipeline | beam.ReadFromGenerator(lambda: load_arrow_files(CACHE_DIR)) | beam.Map(prepare_oscar2301_for_c4) ) # 后续处理流程保持不变 pages = preprocessed deduped_pages = c4_utils.remove_duplicate_text(pages) encoded = deduped_pages | beam.Map(lambda page: json.dumps(page.__dict__, ensure_ascii=False).encode('utf-8')) encoded | 'Write to JSONL' >> beam.io.WriteToText('/content/output.jsonl') pipeline.run()
这个方案适合超大规模数据集,直接从文件系统读取数据,内存占用仅为单个Arrow文件的批次大小。
额外优化建议
- 关闭Dataset的内存缓存:如果之前开启了缓存,执行
dataset['train'].unload()释放已加载的批次数据。 - 调整Flink内存参数:本地运行FlinkRunner时,可通过
pipeline_args设置堆内存上限,避免Flink自身的内存瓶颈。
内容的提问来源于stack exchange,提问作者Mohamed Taher Alrefaie
相关产品推荐
相关产品推荐

