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

训练改进说明

问题诊断

即使有1000张负样本,模型仍然会把无浒苔图片和陆地/海岛预测为浒苔。这说明:

  1. 类别不平衡问题:每张图片中背景像素通常远多于浒苔像素
  2. 损失函数问题:原始损失函数没有给浒苔类别足够的权重
  3. 模型可能学习到"总是预测背景"的策略

改进方案

1. 改进损失函数

1.1 Focal Loss改进

  • 之前alpha=1,所有类别权重相同
  • 现在alpha=[1.0, 3.0],给浒苔类别3倍权重
  • 添加:类别权重tensor,在交叉熵中应用

1.2 Dice Loss改进

  • 之前:使用mean(),所有类别权重相同
  • 现在:对每个类别计算dice loss,然后应用类别权重加权平均

1.3 类别权重

  • 背景权重:1.0(默认)
  • 前景权重:3.0(可调整,建议2.0-5.0)

2. 检查训练数据

运行 check_training_data.py 来:

  • 统计类别分布
  • 计算建议的类别权重
  • 检查数据平衡情况
python check_training_data.py

使用方法

步骤1:检查训练数据

python check_training_data.py

这会输出:

  • 类别分布统计
  • 建议的类别权重
  • 推荐的训练配置

步骤2:更新训练配置

根据检查结果,更新 train_config.json 或使用 train_config_improved.json

{
  "focal_alpha": [1.0, 3.0],
  "background_weight": 1.0,
  "foreground_weight": 3.0,
  ...
}

步骤3:重新训练

python train_seaweed_segmentation.py --config train_config_improved.json

参数说明

focal_alpha

  • 类型:列表 [背景权重, 前景权重]
  • 默认[1.0, 3.0]
  • 说明:Focal Loss的alpha参数,给前景更高的权重

background_weight / foreground_weight

  • 类型:浮点数
  • 默认1.0 / 3.0
  • 说明:类别权重,用于交叉熵和Dice Loss

如何选择权重

  1. 运行数据检查脚本,查看实际的类别分布
  2. 根据建议调整权重
    • 如果浒苔像素 < 1%:使用 foreground_weight >= 5.0
    • 如果浒苔像素 1-5%:使用 foreground_weight = 3.0-5.0
    • 如果浒苔像素 > 5%:使用 foreground_weight = 2.0-3.0

预期效果

改进后的损失函数应该能够:

  1. 更好地学习浒苔特征:给浒苔更高的权重,模型会更关注浒苔区域
  2. 减少误判:模型不会简单地"总是预测背景"
  3. 提高精度:在保持召回率的同时,减少假阳性

进一步优化

如果问题仍然存在:

1. 增加前景权重

"foreground_weight": 5.0,
"focal_alpha": [1.0, 5.0]

2. 调整损失函数权重

"dice_weight": 0.7,  // 增加Dice Loss权重
"focal_weight": 1.0

3. 使用难例挖掘

  • 重点关注模型误判的样本
  • 增加这些样本在训练中的权重

4. 检查数据质量

  • 确保标签正确
  • 检查是否有误标注
  • 确保负样本(无浒苔图片)的标签是全0

验证方法

训练后,检查:

  1. 训练损失:应该逐渐下降
  2. 验证IoU:应该逐渐上升
  3. 预测结果:无浒苔图片应该被正确预测为背景
  4. 混淆矩阵:假阳性(FP)应该减少

注意事项

  1. 权重不要过大:过大的权重可能导致训练不稳定
  2. 监控训练过程:如果损失不下降,可能需要调整权重
  3. 验证集评估:使用验证集来评估改进效果