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

Mask2Former自定义数据集训练异常及ignore_index参数疑问

Mask2Former微调自定义数据集问题排查

一、预测结果全为text类的原因及修复

核心问题点

  1. 预训练模型分类头适配问题
    你使用的facebook/mask2former-swin-small-cityscapes-semantic预训练模型针对Cityscapes的19类语义分割设计,直接通过ignore_mismatched_sizes=True加载时,分类头的权重只是简单截断/扩展,并没有针对你的2类任务重新初始化,导致模型初始就偏向某一类。

  2. 类别标签与模型输入不匹配
    检查训练时传入的class_labels:Mask2Former要求class_labels是每个mask对应的类别ID,如果你的class_labels批量全为1(text类),模型只会学习预测text;另外,若训练集本身存在严重类别失衡(比如text类占比极高),模型会为了降低loss直接输出占比最高的类。

  3. 学习率与优化器选择不当
    Adam优化器配5e-5的学习率对于微调预训练模型来说过小,参数更新幅度不足以纠正初始偏向,导致loss看似在波动但模型没有真正学到分类能力。

  4. 预处理与模型的ignore_index不同步
    你在Mask2FormerImageProcessor中设置了ignore_index=0,但模型初始化时未同步该参数,导致训练时background类的loss被处理器忽略,模型只优化text类的loss,最终只会输出text。

修复方案

  • 重新初始化模型分类头
    不要依赖自动适配,手动修改配置后加载模型,确保分类头针对2类任务初始化:

    from transformers import Mask2FormerConfig
    id2label = {0:"background",1:"text"}
    label2id = {v:k for k,v in id2label.items()}
    
    config = Mask2FormerConfig.from_pretrained(
        "facebook/mask2former-swin-small-cityscapes-semantic",
        num_labels=len(id2label),
        id2label=id2label,
        label2id=label2id,
        ignore_index=0
    )
    model = Mask2FormerForUniversalSegmentation.from_pretrained(
        "facebook/mask2former-swin-small-cityscapes-semantic",
        config=config,
        ignore_mismatched_sizes=True
    )
    
  • 调整优化器与学习率
    改用AdamW优化器(更适合预训练模型微调),将学习率提升至1e-4:

    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-5)
    
  • 验证训练集标签与输入
    打印batch["class_labels"]和batch["mask_labels"]的唯一值,确认训练集中存在background类(0类)的样本,且标签输入正确。

二、ignore_index设为0时的警告与mIoU=1.0问题

原因解析

  • 警告触发逻辑:当设置ignore_index=0时,计算IoU会忽略background类,仅计算text类的IoU。如果测试集中没有text类样本,或者预测与标签的text类区域无交集,会出现total_area_union=0的情况,触发除以0的警告。
  • mIoU=1.0的本质:当忽略0类后,若测试集标签全为background(0类),此时没有有效计算样本,metric会默认返回1.0;或者模型全预测text,而测试集中text类的标签区域与预测区域完全重叠(这种情况概率极低)。

处理建议

  • 保留ignore_index=255:对于二分类语义分割,应计算全类别的mIoU,不要忽略background类,除非你的任务仅关注text类的分割效果。
  • 单独提取目标类IoU:如果只关心text类的表现,直接从metric结果中提取对应类别的IoU:
    result = metric.compute(num_labels=len(id2label), reduce_labels=False, ignore_index=255)
    print("Text类IoU:", result['per_category_iou'][1])
    
  • 检查测试集标签:确认测试集中存在text类(1类)的样本,避免因无有效计算样本导致的异常结果。

内容的提问来源于stack exchange,提问作者Phourior

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 18:35:37