You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用本地数据集微调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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.09 05:09:56