使用ViTFeatureExtractor执行with_transform后数据集结构异常及操作失败问题求助
使用ViTFeatureExtractor执行with_transform后数据集结构异常及操作失败问题求助
我在使用ViT的特征提取器时,遇到了一些无法理解的奇怪问题,想请大家帮忙解惑。
加载数据集后,我查看训练集的特征结构如下:
ds['train'].features {'image_file_path': Value(dtype='string', id=None), 'image': Image(mode=None, decode=True, id=None), 'labels': ClassLabel(names=['angular_leaf_spot', 'bean_rust', 'healthy'], id=None)}
这时候不管是按列取标签,还是按行取样本,都能正常访问:
# 按列取前5个标签 ds['train']['labels'][0:5] # 输出:[0, 0, 0, 0, 0] # 取前2个样本 ds['train'][0:2] # 输出: {'image_file_path': ['/home/albert/.cache/huggingface/datasets/downloads/extracted/967f0d9f61a7a8de58892c6fab6f02317c06faf3e19fba6a07b0885a9a7142c7/train/angular_leaf_spot/angular_leaf_spot_train.0.jpg', '/home/albert/.cache/huggingface/datasets/downloads/extracted/967f0d9f61a7a8de58892c6fab6f02317c06faf3e19fba6a07b0885a9a7142c7/train/angular_leaf_spot/angular_leaf_spot_train.1.jpg'], 'image': [<PIL.JpegImagePlugin.JpegImageFile image mode=RGB size=500x500>, <PIL.JpegImagePlugin.JpegImageFile image mode=RGB size=500x500>], 'labels': [0, 0]}
但是当我用ViTFeatureExtractor定义转换函数,并用with_transform处理数据集后:
from transformers import ViTFeatureExtractor model_name_or_path = 'google/vit-base-patch16-224-in21k' feature_extractor = ViTFeatureExtractor.from_pretrained(model_name_or_path) ds = load_dataset('beans') def transform(example_batch): inputs = feature_extractor([x for x in example_batch['image']], return_tensors='pt') inputs['labels'] = example_batch['labels'] return inputs prepared_ds = ds.with_transform(transform)
查看prepared_ds['train'].features,显示的特征还是和原数据集一致,但访问样本时,返回的是处理后的pixel_values和标签:
prepared_ds['train'][0:2] # 输出: {'pixel_values': tensor([[[[-0.5686, -0.5686, -0.5608, ..., -0.0275, 0.1843, -0.2471], ..., [-0.5843, -0.5922, -0.6078, ..., 0.2627, 0.1608, 0.2000]], [[-0.7098, -0.7098, -0.7490, ..., -0.3725, -0.1608, -0.6000], ..., [-0.8824, -0.9059, -0.9216, ..., -0.2549, -0.2000, -0.1216]]], [[[-0.5137, -0.4902, -0.4196, ..., -0.0275, -0.0039, -0.2157], ..., [-0.5216, -0.5373, -0.5451, ..., -0.1294, -0.1529, -0.2627]], [[-0.1843, -0.2000, -0.1529, ..., 0.2157, 0.2078, -0.0902], ..., [-0.7725, -0.7961, -0.8039, ..., -0.3725, -0.4196, -0.5451]], [[-0.7569, -0.8510, -0.8353, ..., -0.3255, -0.2706, -0.5608], ..., [-0.5294, -0.5529, -0.5608, ..., -0.1686, -0.1922, -0.3333]]]]), 'labels': [0, 0]}
问题1:按列访问标签触发KeyError
当我尝试直接按列访问标签时,执行prepared_ds['train']['labels'],却触发了KeyError,错误栈如下:
--------------------------------------------------------------------------- KeyError Traceback (most recent call last) Cell In[32], line 1 ----> 1 prepared_ds['train']['labels'] File ~/anaconda3/envs/LLM/lib/python3.12/site-packages/datasets/arrow_dataset.py:2872, in Dataset.__getitem__(self, key) 2870 def __getitem__(self, key): # noqa: F811 2871 """Can be used to index columns (by string names) or rows (by integer index or iterable of indices or bools).""" -> 2872 return self._getitem(key) File ~/anaconda3/envs/LLM/lib/python3.12/site-packages/datasets/arrow_dataset.py:2857, in Dataset._getitem(self, key, **kwargs) 2855 formatter = get_formatter(format_type, features=self._info.features, **format_kwargs) 2856 pa_subtable = query_table(self._data, key, indices=self._indices) -> 2857 formatted_output = format_table( 2858 pa_subtable, key, formatter=formatter, format_columns=format_columns, output_all_columns=output_all_columns 2859 ) 2860 return formatted_output File ~/anaconda3/envs/LLM/lib/python3.12/site-packages/datasets/formatting/formatting.py:639, in format_table(table, key, formatter, format_columns, output_all_columns) 637 python_formatter = PythonFormatter(features=formatter.features) 638 if format_columns is None: --> 639 return formatter(pa_table, query_type=query_type) 640 elif query_type == "column": 641 if key in format_columns: File ~/anaconda3/envs/LLM/lib/python3.12/site-packages/datasets/formatting/formatting.py:405, in Formatter.__call__(self, pa_table, query_type) 403 return self.format_row(pa_table) 404 elif query_type == "column": --> 405 return self.format_column(pa_table) 406 elif query_type == "batch": 407 return self.format_batch(pa_table) File ~/anaconda3/envs/LLM/lib/python3.12/site-packages/datasets/formatting/formatting.py:501, in CustomFormatter.format_column(self, pa_table) 500 def format_column(self, pa_table: pa.Table) -> ColumnFormat: --> 501 formatted_batch = self.format_batch(pa_table) 502 if hasattr(formatted_batch, "keys"): 503 if len(formatted_batch.keys()) > 1: File ~/anaconda3/envs/LLM/lib/python3.12/site-packages/datasets/formatting/formatting.py:522, in CustomFormatter.format_batch(self, pa_table) 520 batch = self.python_arrow_extractor().extract_batch(pa_table) 521 batch = self.python_features_decoder.decode_batch(batch) --> 522 return self.transform(batch) Cell In[12], line 5, in transform(example_batch) 3 def transform(example_batch): 4 # Take a list of PIL images and turn them to pixel values --> 5 inputs = feature_extractor([x for x in example_batch['image']], return_tensors='pt') 7 # Don't forget to include the labels! 8 inputs['labels'] = example_batch['labels'] KeyError: 'image'
看起来错误是因为当我按列访问时,转换函数被触发了,但此时的example_batch里没有image字段,这让我很困惑——为什么只是访问标签,会重新执行转换函数?
问题2:保存处理后的数据集触发TypeError
另外,当我尝试把处理后的数据集保存到磁盘时,执行prepared_ds.save_to_disk(img_path),又遇到了TypeError:
--------------------------------------------------------------------------- TypeError Traceback (most recent call last) Cell In[21], line 1 ----> 1 dataset.save_to_disk(img_path) File ~/anaconda3/envs/LLM/lib/python3.13/site-packages/datasets/arrow_dataset.py:1503, in Dataset.save_to_disk(self, dataset_path, max_shard_size, num_shards, num_proc, storage_options) 1501 json.dumps(state["_format_kwargs"][k]) 1502 except TypeError as e: -> 1503 raise TypeError( 1504 str(e) + f"\nThe format kwargs must be JSON serializable, but key '{k}' isn't." 1505 ) from None 1506 # Get json serializable dataset info 1507 dataset_info = asdict(self._info) TypeError: Object of type function is not JSON serializable The format kwargs must be JSON serializable, but key 'transform' isn't.
我的疑问
需要说明的是,原示例里的训练、评估等流程都是正常工作的,我只是在尝试探索数据集结构、保存处理后的数据集时才遇到这些问题。
我想请教大家:
- 为什么执行
with_transform()或set_transform()后,数据集的访问方式不能和原来保持一致? - 为什么只是尝试访问某个特征,会触发转换函数的执行?
- 有没有办法在应用转换后,还能正常访问原有的特征,或者顺利保存数据集?
希望有人能帮我理清这个行为背后的原因,谢谢!
备注:内容来源于stack exchange,提问作者hamagust
相关产品推荐
相关产品推荐

