如何解决Huggingface Dataset.map合并分片时的OOM问题?
问题描述
使用Hugging Face Dataset.map 函数执行以下代码时:
dataset.map(myfunc, num_proc=16, keep_in_memory=False, cache_file_name='parts.arrow', batch_size=16, writer_batch_size=16 )
因数据集规模过大触发内存不足(OOM),报错信息如下:
/site-packages/datasets/table.py:1421: table = cls._concat_blocks(blocks, axis=0) Killed
观察到map执行到进度条100%时,分片Arrow文件已生成,但在合并分片阶段触发内存耗尽。单分片Arrow文件大小约100+GB,机器内存为80GB,寻求解决_concat_blocks函数OOM问题的方案。
解决方法
跳过自动合并,手动处理分片
map默认会合并多进程生成的分片,可通过load_from_cache_file=False跳过该步骤,直接获取分片数据集列表,后续按需逐个处理,避免一次性加载所有分片:# 执行map时跳过自动合并,得到分片数据集列表 transformed_shards = dataset.map( myfunc, num_proc=16, keep_in_memory=False, cache_file_name='parts.arrow', batch_size=16, writer_batch_size=16, load_from_cache_file=False # 关键参数:禁用自动合并 ) # 逐个加载分片处理,降低内存占用 for shard in transformed_shards: process_shard(shard) # 替换为你的分片处理逻辑调大
writer_batch_size减少分片数量
当前writer_batch_size=16会生成大量小分片,合并时需同时加载多个分片到内存。调大该参数可减少分片总数,降低合并阶段的内存压力:dataset.map( myfunc, num_proc=16, keep_in_memory=False, cache_file_name='parts.arrow', batch_size=16, writer_batch_size=1000 # 增大值以减少分片数量 )分批手动合并分片
若必须合并为单一Dataset,可采用分批合并的方式,避免一次性加载所有分片:from datasets import concatenate_datasets, load_from_disk transformed_shards = dataset.map( myfunc, num_proc=16, keep_in_memory=False, cache_file_name='parts.arrow', batch_size=16, writer_batch_size=16, load_from_cache_file=False ) # 每次合并2个分片,可根据内存情况调整批次大小 batch_merge_size = 2 temp_datasets = [] for i in range(0, len(transformed_shards), batch_merge_size): batch_shards = transformed_shards[i:i+batch_merge_size] merged_temp = concatenate_datasets(batch_shards) merged_temp.save_to_disk(f"temp_merged_{i//batch_merge_size}") temp_datasets.append(f"temp_merged_{i//batch_merge_size}") merged_temp = None # 释放内存 # 合并所有临时文件得到最终数据集 final_shards = [load_from_disk(path) for path in temp_datasets] final_dataset = concatenate_datasets(final_shards)使用流式处理替代全量合并
若无需完整Dataset对象,可直接以流式方式遍历分片文件,完全跳过合并操作:import pyarrow.dataset as pad # 读取所有分片Arrow文件 arrow_dataset = pad.dataset("parts.arrow", format="arrow") # 流式遍历数据批次 for batch in arrow_dataset.to_batches(): process_batch(batch.to_pandas()) # 替换为你的批次处理逻辑升级Datasets版本
旧版本的合并逻辑可能存在内存效率缺陷,升级到最新版Hugging Face Datasets可能修复该问题:pip install --upgrade datasets
内容的提问来源于stack exchange,提问作者alvas
相关产品推荐
相关产品推荐

