加载GAIR/MathPile数据集时遇类型转换错误求助
解决GAIR/MathPile数据集多进程加载时的PyArrow类型转换错误
问题分析
报错TypeError: Couldn't cast array of type string to null本质是多进程加载时,不同进程推断出的数据集Schema不一致(部分文件中某列被推断为string,另一部分被推断为null),合并阶段触发类型冲突。
解决步骤
1. 单进程验证数据集完整性
先禁用多进程加载,确认数据集本身无损坏或格式异常:
from datasets import Dataset, load_dataset import os def get_hf_dataset_gair(path: str = '~/data/GAIR/MathPile/train/') -> Dataset: path: str = os.path.expanduser(path) # 单进程加载,排查核心问题 dataset = load_dataset(path, split='train') print(dataset[0]) print(dataset.schema) # 查看自动推断的字段类型 # 后续处理逻辑不变 all_columns = dataset.column_names all_columns.remove('text') dataset = dataset.remove_columns(all_columns) dataset = dataset.shuffle(seed=42) dataset = dataset.select(10_000) return dataset get_hf_dataset_gair()
如果单进程能成功运行,说明问题出在多进程下的Schema合并;如果仍报错,需检查数据集文件是否损坏。
2. 手动指定强制Schema
明确定义数据集Schema,避免自动推断的不一致:
from datasets import Dataset, load_dataset import os import pyarrow as pa def get_hf_dataset_gair(path: str = '~/data/GAIR/MathPile/train/') -> Dataset: path: str = os.path.expanduser(path) # 根据单进程加载的schema,定义固定字段类型 schema = pa.schema([ ('text', pa.string()), ('meta', pa.string()), # 替换为实际存在的其他字段及对应类型 ]) # 加载时指定schema,开启多进程 dataset = load_dataset(path, split='train', num_proc=os.cpu_count(), schema=schema) print(dataset[0]) all_columns = dataset.column_names all_columns.remove('text') dataset = dataset.remove_columns(all_columns) dataset = dataset.shuffle(seed=42) dataset = dataset.select(10_000) return dataset get_hf_dataset_gair()
3. 检查解压后的数据集文件
- 确认所有
.gz文件已正确解压:
cd ~/data/GAIR/MathPile/train/ ls | grep .gz # 无输出则说明全部解压完成
- 若存在未解压文件,重新执行解压:
find . -type f -name "*.gz" -exec gzip -d {} \;
4. 调整datasets库版本
2.x.x部分子版本存在多进程Schema合并bug,尝试切换到稳定版本:
# 升级到最新稳定版 pip install --upgrade datasets==2.18.0 # 或降级到已知兼容版本 pip install datasets==2.14.0
5. 先单进程加载再并行处理
如果多进程加载的Schema冲突无法避免,可先单进程加载数据集,后续处理步骤开启多进程:
from datasets import Dataset, load_dataset import os def get_hf_dataset_gair(path: str = '~/data/GAIR/MathPile/train/') -> Dataset: path: str = os.path.expanduser(path) # 单进程加载数据集 dataset = load_dataset(path, split='train') # 移除冗余列时开启多进程 dataset = dataset.remove_columns([col for col in dataset.column_names if col != 'text'], num_proc=os.cpu_count()) dataset = dataset.shuffle(seed=42) dataset = dataset.select(10_000) return dataset get_hf_dataset_gair()
内容的提问来源于stack exchange,提问作者Charlie Parker
相关产品推荐
相关产品推荐

