使用Transformers的image_utils遇“图像维度不支持”错误的排查与解决
微调ViT时“Unsupported number of image dimensions”错误的调试与解决
我按照HuggingFace的ViT微调教程操作,使用官方beans数据集一切正常,但换成自定义数据集后,触发了ValueError: Unsupported number of image dimensions: 2错误。
报错栈信息
--------------------------------------------------------------------------- ValueError Traceback (most recent call last) /tmp/ipykernel_2042949/883871373.py in <module> ----> 1 train_results = trainer.train() 2 trainer.save_model() 3 trainer.log_metrics("train", train_results.metrics) 4 trainer.save_metrics("train", train_results.metrics) 5 trainer.save_state() ~/miniconda3/lib/python3.9/site-packages/transformers/trainer.py in train(self, resume_from_checkpoint, trial, ignore_keys_for_eval, **kwargs) 1532 self._inner_training_loop, self._train_batch_size, args.auto_find_batch_size 1533 ) -> 1534 return inner_training_loop( 1535 args=args, 1536 resume_from_checkpoint=resume_from_checkpoint, ~/miniconda3/lib/python3.9/site-packages/transformers/trainer.py in _inner_training_loop(self, batch_size, args, resume_from_checkpoint, trial, ignore_keys_for_eval) 1754 1755 step = -1 -> 1756 for step, inputs in enumerate(epoch_iterator): 1757 1758 # Skip past any already trained steps if resuming training ~/miniconda3/lib/python3.9/site-packages/torch/utils/data/dataloader.py in __next__(self) 626 # TODO(https://github.com/pytorch/pytorch/issues/76750) ... --> 119 raise ValueError(f"Unsupported number of image dimensions: {image.ndim}") 120 121 if image.shape[first_dim] in (1, 3): ValueError: Unsupported number of image dimensions: 2
错误来自transformers库的image_utils.py文件。
调试过程
- 对比自定义数据集与beans数据集的张量形状,确认二者一致:
$ prepared_ds['train'][0:2]['pixel_values'].shape torch.Size([2, 3, 224, 224])
- 根据报错栈定位到
infer_channel_dimension_format函数,编写代码定位问题图片:
from transformers.image_utils import infer_channel_dimension_format try: for i, img in enumerate(prepared_ds["train"]): infer_channel_dimension_format(img["pixel_values"]) except ValueError as ve: print(i+1)
- 检查定位到的图片,发现其为灰度图(模式为L),而非模型要求的RGB格式:
$ ds["train"][8] {'image': <PIL.JpegImagePlugin.JpegImageFile image mode=L size=390x540>, 'image_file_path': '/data/alamy/img/00000/000001069.jpg', 'labels': 0}
解决方案
在数据转换函数中添加convert('RGB'),将所有图片统一转为RGB格式:
def transform(example_batch): # Take a list of PIL images and turn them to pixel values inputs = feature_extractor([x.convert("RGB") for x in example_batch['image']], return_tensors='pt') # Don't forget to include the labels! inputs['labels'] = example_batch['labels'] return inputs
内容的提问来源于stack exchange,提问作者Pablo Mendes
相关产品推荐
相关产品推荐

