cuibinge's picture
Sync YOLO training and evaluation utilities (part 2)
c2b1b26 verified
|
Raw
History Blame Contribute Delete
3.44 kB

浒苔预测阈值修复说明

问题描述

当处理整张没有浒苔的256x256图片时,模型可能会把整张图片全部预测为浒苔。这是因为:

  1. 原始预测逻辑:使用 torch.argmax() 直接选择概率最高的类别
  2. 问题原因:当两个类别的概率接近时(比如都是0.5左右),argmax 会选择概率稍高的那个
  3. 类别不平衡:如果训练数据中浒苔样本较多,模型可能倾向于预测浒苔类别

解决方案

添加了概率阈值机制

  • 修改前:直接使用 argmax 选择概率最高的类别
  • 修改后:只有当浒苔(类别1)的概率超过阈值(默认0.5)时,才预测为浒苔

代码修改

predict_image() 函数中:

# 修改前
pred_class = torch.argmax(pred, dim=1).squeeze(0).cpu().numpy()

# 修改后
seaweed_prob = pred[0, 1].squeeze(0).cpu().numpy()
pred_class = (seaweed_prob > threshold).astype(np.uint8)

使用方法

1. 基本使用(使用默认阈值0.5)

直接运行预测脚本即可,默认阈值为0.5:

python predict_seaweed_256x256_simple.py

2. 自定义阈值

在脚本的 main() 函数中修改 prediction_threshold 变量:

# 预测阈值:只有当浒苔概率超过此阈值时才预测为浒苔
# 建议值:0.5-0.7,可以根据验证集调整
prediction_threshold = 0.5  # 修改为你需要的值

3. 阈值选择建议

  • 0.5:平衡阈值,适合大多数情况
  • 0.6-0.7:更保守,减少误报(假阳性),但可能漏检一些浒苔
  • 0.4-0.5:更宽松,减少漏检(假阴性),但可能增加误报

建议:在验证集上测试不同阈值,选择F1分数最高的阈值。

修改的文件

  1. predict_seaweed_256x256_simple.py - 单张256x256图片预测脚本
  2. predict_large_image_tiles_256.py - 大图瓦片预测脚本

验证方法

运行预测后,检查以下内容:

  1. 概率图:查看 *_probability.png 文件,确认无浒苔图片的概率值是否低于阈值
  2. 预测结果:查看 *_prediction.png 文件,确认无浒苔图片是否被正确预测为背景
  3. 统计信息:查看 prediction_results.json,检查无浒苔图片的 seaweed_ratio 是否接近0

进一步优化建议

如果问题仍然存在,可以考虑:

  1. 调整阈值:根据验证集结果调整阈值
  2. 重新训练模型:增加无浒苔样本的训练数据
  3. 使用类别权重:在训练时给背景类别更高的权重
  4. 后处理:添加形态学操作(如开运算)去除小的误检区域

技术细节

为什么会出现这个问题?

  1. Softmax特性:softmax会将logits转换为概率分布,即使两个类别的logits很接近,softmax后的概率也会强制归一化
  2. 模型倾向性:如果训练数据中浒苔样本较多,模型可能学习到倾向于预测浒苔的模式
  3. 边界情况:当输入特征不明显时,模型可能给出不确定的预测

阈值机制的工作原理

  • 阈值机制相当于在softmax概率上添加了一个"置信度门控"
  • 只有当模型对浒苔的预测足够确信(概率>阈值)时,才预测为浒苔
  • 这样可以有效减少在不确定情况下的误判