训练改进说明
问题诊断
即使有1000张负样本,模型仍然会把无浒苔图片和陆地/海岛预测为浒苔。这说明:
- 类别不平衡问题:每张图片中背景像素通常远多于浒苔像素
- 损失函数问题:原始损失函数没有给浒苔类别足够的权重
- 模型可能学习到"总是预测背景"的策略
改进方案
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%:使用
foreground_weight >= 5.0 - 如果浒苔像素 1-5%:使用
foreground_weight = 3.0-5.0 - 如果浒苔像素 > 5%:使用
foreground_weight = 2.0-3.0
- 如果浒苔像素 < 1%:使用
预期效果
改进后的损失函数应该能够:
- 更好地学习浒苔特征:给浒苔更高的权重,模型会更关注浒苔区域
- 减少误判:模型不会简单地"总是预测背景"
- 提高精度:在保持召回率的同时,减少假阳性
进一步优化
如果问题仍然存在:
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
验证方法
训练后,检查:
- 训练损失:应该逐渐下降
- 验证IoU:应该逐渐上升
- 预测结果:无浒苔图片应该被正确预测为背景
- 混淆矩阵:假阳性(FP)应该减少
注意事项
- 权重不要过大:过大的权重可能导致训练不稳定
- 监控训练过程:如果损失不下降,可能需要调整权重
- 验证集评估:使用验证集来评估改进效果