如何修复interleave_datasets仅采样单个数据集的异常问题?
问题描述
使用Hugging Face Datasets库的interleave_datasets函数,按[0.5, 0.5]的概率对流式加载的c4英文训练集和wikitext-103-v1训练集做交错处理,但实际采样时连续100个样本全来自c4数据集(样本字段为text、timestamp、url),完全没出现wikitext的样本。
附代码
from datasets import load_dataset from datasets import interleave_datasets # 加载数据集 c4 = load_dataset("c4", "en", split="train", streaming=True) wikitext = load_dataset("wikitext", "wikitext-103-v1", split="train", streaming=True) # 交错处理数据集 datasets = [c4, wikitext] for dataset in datasets: print(dataset.description) interleaved = interleave_datasets(datasets, probabilities=[0.5, 0.5]) print(interleaved)
采样输出示例
example.keys()=dict_keys(['text', 'timestamp', 'url']) example.keys()=dict_keys(['text', 'timestamp', 'url']) ...(重复至100次) counts=100
异常原因
- 数据集字段不兼容:c4样本有
text、timestamp、url三个字段,而wikitext-103-v1的样本只有text字段。interleave_datasets处理流式数据集时,会自动对齐字段,默认丢弃结构不匹配的样本,导致wikitext的所有样本被过滤,只剩c4的样本。 - 流式模式的静默过滤:流式加载下,库不会抛出字段不匹配的报错,而是直接静默过滤不符合结构的样本,所以看起来只有c4的数据被采样到。
修复方案
核心是统一两个数据集的字段结构,确保样本字段完全一致。下面以给wikitext补全缺失字段为例:
修复后代码
from datasets import load_dataset from datasets import interleave_datasets # 加载c4数据集 c4 = load_dataset("c4", "en", split="train", streaming=True) # 给wikitext添加缺失字段,值设为None def add_missing_fields(example): example['timestamp'] = None example['url'] = None return example wikitext = load_dataset("wikitext", "wikitext-103-v1", split="train", streaming=True).map(add_missing_fields) # 重新执行交错处理 datasets = [c4, wikitext] interleaved = interleave_datasets(datasets, probabilities=[0.5, 0.5]) # 验证采样结果 sample_counts = {'c4': 0, 'wikitext': 0} for example in interleaved.take(100): # 通过url是否非空判断样本来源 if example['url'] is not None: sample_counts['c4'] += 1 else: sample_counts['wikitext'] += 1 print(sample_counts)
说明
- 通过
map函数为wikitext补全c4的额外字段,消除结构差异,避免被过滤。 - 验证时可以通过字段值的区别(比如wikitext的
url为None)区分样本来源,确认两个数据集的样本都能被正常采样。
内容的提问来源于stack exchange,提问作者Charlie Parker
相关产品推荐
相关产品推荐

