使用Huggingface datasets时n转m样本映射触发ArrowInvalid错误
Huggingface Datasets批量映射时样本数不匹配触发pyarrow.lib.ArrowInvalid错误
问题场景
使用Huggingface Datasets库(版本2.5.2)时,按照官方文档说明,当batched=True且batch_size>1时,映射函数可接收n个样本的批次并返回任意数量的样本批次。但在加载本地图像并剔除损坏图像(输入样本数n大于输出样本数m)的场景中,触发了pyarrow.lib.ArrowInvalid错误。
复现代码
from datasets import Dataset import pandas as pd dataset_ = Dataset.from_pandas( pd.DataFrame({ 'path': ['doc1.jpg', 'doc2.jpg', 'doc3.jpg'], 'documentType': [1,2,3] }) ) def fun_map(examples): return {"output": [1,2]} preprocessed_dataset = dataset_.map(fun_map, batched=True, batch_size=2, num_proc=2)
报错信息
#0: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 1/1 [00:00<00:00, 328.40ba/s] #1: 0%| | 0/1 [00:00<?, ?ba/s] multiprocess.pool.RemoteTraceback: | 0/1 [00:00<?, ?ba/s] """ Traceback (most recent call last): File "/opt/conda/lib/python3.8/site-packages/multiprocess/pool.py", line 125, in worker result = (True, func(*args, **kwds)) File "/opt/conda/lib/python3.8/site-packages/datasets/arrow_dataset.py", line 578, in wrapper out: Union["Dataset", "DatasetDict"] = func(self, *args, **kwargs) File "/opt/conda/lib/python3.8/site-packages/datasets/arrow_dataset.py", line 545, in wrapper out: Union["Dataset", "DatasetDict"] = func(self, *args, **kwargs) File "/opt/conda/lib/python3.8/site-packages/datasets/fingerprint.py", line 480, in wrapper out = func(self, *args, **kwargs) File "/opt/conda/lib/python3.8/site-packages/datasets/arrow_dataset.py", line 2867, in _map_single writer.write_batch(batch) File "/opt/conda/lib/python3.8/site-packages/datasets/arrow_writer.py", line 527, in write_batch pa_table = pa.Table.from_arrays(arrays, schema=schema) File "pyarrow/table.pxi", line 3597, in pyarrow.lib.Table.from_arrays File "pyarrow/table.pxi", line 2793, in pyarrow.lib.Table.validate File "pyarrow/error.pxi", line 100, in pyarrow.lib.check_status pyarrow.lib.ArrowInvalid: Column 2 named output expected length 1 but got length 2 """ The above exception was the direct cause of the following exception: Traceback (most recent call last): File "<stdin>", line 1, in <module> File "/opt/conda/lib/python3.8/site-packages/datasets/arrow_dataset.py", line 2563, in map transformed_shards[index] = async_result.get() File "/opt/conda/lib/python3.8/site-packages/multiprocess/pool.py", line 771, in get raise self._value pyarrow.lib.ArrowInvalid: Column 2 named output expected length 1 but got length 2
解决方案
在map方法中添加remove_columns参数,指定移除原数据集的所有列,让输出的新列直接替代原数据集结构,避免新旧列样本数不匹配的冲突。
修改后的代码
preprocessed_dataset = dataset_.map(fun_map, batched=True, batch_size=2, num_proc=2, remove_columns=dataset_.column_names)
内容的提问来源于stack exchange,提问作者samusa
相关产品推荐
相关产品推荐

