使用本地数据集微调SegFormer遇TypeError:无法处理PngImageFile
本地数据集微调SegFormer语义分割:报错解决与实现流程
问题场景
首次使用HuggingFace,用本地PNG格式图像和标签构建数据集,尝试微调SegFormer时触发TypeError: must be real number, not PngImageFile错误。原始代码如下:
数据集创建代码
def create_dataset(image_paths, label_paths): dataset = datasets.Dataset.from_dict({"image": sorted(image_paths), "label": sorted(label_paths)}) dataset = dataset.cast_column("image", datasets.Image()) dataset = dataset.cast_column("label", datasets.Image()) return dataset
模型加载与训练代码
model = SegformerForSemanticSegmentation.from_pretrained("nvidia/segformer-b0-finetuned-ade-512-512") training_args = TrainingArguments(output_dir="test_trainer") trainer = Trainer( model=model, args=training_args, train_dataset=train_dataset, eval_dataset=validation_dataset, compute_metrics=compute_metrics, ) trainer.train()
完整报错信息
0%| | 0/93 [00:00<?, ?it/s]Traceback (most recent call last): File "C:\Users\itishel\PycharmProjects\huggingSegFormer\lib\site-packages\accelerate\data_loader.py", line 384, in __iter__ current_batch = next(dataloader_iter) File "C:\Users\itishel\PycharmProjects\huggingSegFormer\lib\site-packages\torch\utils\data\dataloader.py", line 681, in __next__ data = self._next_data() File "C:\Users\itishel\PycharmProjects\huggingSegFormer\lib\site-packages\torch\utils\data\dataloader.py", line 721, in _next_data data = self._dataset_fetcher.fetch(index) # may raise StopIteration File "C:\Users\itishel\PycharmProjects\huggingSegFormer\lib\site-packages\torch\utils\data_utils\fetch.py", line 52, in fetch return self.collate_fn(data) File "C:\Users\itishel\PycharmProjects\huggingSegFormer\lib\site-packages\transformers\data\data_collator.py", line 70, in default_data_collator return torch_default_data_collator(features) File "C:\Users\itishel\PycharmProjects\huggingSegFormer\lib\site-packages\transformers\data\data_collator.py", line 119, in torch_default_data_collator batch["labels"] = torch.tensor([f["label"] for f in features], dtype=dtype) TypeError: must be real number, not PngImageFile
报错核心原因
用datasets.Image()转换标签列后,每条数据的label是PngImageFile对象,而默认数据拼接器torch_default_data_collator会尝试直接将其转为张量,但该对象并非数值类型,导致转换失败。SegFormer要求标签是单通道整数张量(每个像素对应类别ID),而非图像对象。
分步解决方案
1. 重构数据集处理逻辑,将标签转为类别ID张量
放弃直接用cast_column将标签转为Image类型,改用自定义函数加载标签并转换为符合要求的张量:
from PIL import Image import numpy as np import torch import datasets def process_example(example): # 加载图像并转为RGB格式 example["image"] = Image.open(example["image"]).convert("RGB") # 加载标签:若为灰度标签(单通道)直接转为整数张量 label = Image.open(example["label"]).convert("L") example["label"] = torch.tensor(np.array(label), dtype=torch.long) return example def create_dataset(image_paths, label_paths): dataset = datasets.Dataset.from_dict({ "image": sorted(image_paths), "label": sorted(label_paths) }) # 应用自定义处理函数 dataset = dataset.map(process_example) # 设置数据集格式为PyTorch张量,方便DataLoader处理 dataset.set_format("torch", columns=["image", "label"]) return dataset
如果是伪彩色标签(RGB格式),需要先将颜色映射为类别ID:
# 替换为你的颜色-类别ID映射关系 color_to_id = {(0, 0, 0): 0, (255, 255, 255): 1, (128, 0, 0): 2} def process_pseudocolor_label(example): example["image"] = Image.open(example["image"]).convert("RGB") label_img = Image.open(example["label"]).convert("RGB") label_np = np.array(label_img) # 逐像素映射到类别ID label_id = np.zeros((label_np.shape[0], label_np.shape[1]), dtype=np.int64) for color, id in color_to_id.items(): mask = np.all(label_np == color, axis=-1) label_id[mask] = id example["label"] = torch.tensor(label_id, dtype=torch.long) return example
2. 添加图像预处理管道
SegFormer需要输入图像符合指定尺寸和归一化标准,加入预处理和可选的数据增强:
from torchvision import transforms # 训练集预处理(含数据增强) train_transforms = transforms.Compose([ transforms.Resize((512, 512)), # 匹配SegFormer-b0的输入尺寸 transforms.RandomHorizontalFlip(p=0.5), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 验证集预处理(无增强) val_transforms = transforms.Compose([ transforms.Resize((512, 512)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 修改处理函数,加入预处理 def process_train_example(example): example["image"] = train_transforms(Image.open(example["image"]).convert("RGB")) label = Image.open(example["label"]).convert("L") # 标签用最近邻插值,避免类别ID模糊 label = transforms.Resize((512, 512), interpolation=transforms.InterpolationMode.NEAREST)(label) example["label"] = torch.tensor(np.array(label), dtype=torch.long) return example def process_val_example(example): example["image"] = val_transforms(Image.open(example["image"]).convert("RGB")) label = Image.open(example["label"]).convert("L") label = transforms.Resize((512, 512), interpolation=transforms.InterpolationMode.NEAREST)(label) example["label"] = torch.tensor(np.array(label), dtype=torch.long) return example # 分别创建训练/验证集 train_dataset = datasets.Dataset.from_dict({"image": train_img_paths, "label": train_label_paths}).map(process_train_example) validation_dataset = datasets.Dataset.from_dict({"image": val_img_paths, "label": val_label_paths}).map(process_val_example) train_dataset.set_format("torch", columns=["image", "label"]) validation_dataset.set_format("torch", columns=["image", "label"])
3. 适配模型分类头到自定义类别数
预训练SegFormer针对ADE20K数据集(150类),需替换分类头适配你的数据集类别数:
from transformers import SegformerForSemanticSegmentation num_classes = 2 # 替换为你的数据集实际类别数 model = SegformerForSemanticSegmentation.from_pretrained( "nvidia/segformer-b0-finetuned-ade-512-512", num_labels=num_classes, ignore_mismatched_sizes=True # 强制替换分类头 )
4. 使用语义分割专用数据拼接器(可选)
避免默认拼接器的兼容问题,使用HuggingFace提供的语义分割专用DataCollator:
from transformers import DataCollatorForSegmentation data_collator = DataCollatorForSegmentation() # 初始化Trainer时传入 trainer = Trainer( model=model, args=training_args, train_dataset=train_dataset, eval_dataset=validation_dataset, compute_metrics=compute_metrics, data_collator=data_collator )
5. 修正评估指标函数
确保compute_metrics正确处理模型输出和标签,示例用mIoU和准确率:
import evaluate metric = evaluate.load("mean_iou") def compute_metrics(eval_pred): logits, labels = eval_pred predictions = np.argmax(logits, axis=1) # 计算指标,ignore_index替换为你的数据集忽略类ID(如无则删除) metrics = metric.compute( predictions=predictions, references=labels, num_labels=num_classes, ignore_index=255 ) return { "mIoU": metrics["mean_iou"], "accuracy": metrics["mean_accuracy"] }
内容的提问来源于stack exchange,提问作者idoo
相关产品推荐
相关产品推荐

