| # 训练改进说明 | |
| ## 问题诊断 | |
| 即使有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` 来: | |
| - 统计类别分布 | |
| - 计算建议的类别权重 | |
| - 检查数据平衡情况 | |
| ```bash | |
| python check_training_data.py | |
| ``` | |
| ## 使用方法 | |
| ### 步骤1:检查训练数据 | |
| ```bash | |
| python check_training_data.py | |
| ``` | |
| 这会输出: | |
| - 类别分布统计 | |
| - 建议的类别权重 | |
| - 推荐的训练配置 | |
| ### 步骤2:更新训练配置 | |
| 根据检查结果,更新 `train_config.json` 或使用 `train_config_improved.json`: | |
| ```json | |
| { | |
| "focal_alpha": [1.0, 3.0], | |
| "background_weight": 1.0, | |
| "foreground_weight": 3.0, | |
| ... | |
| } | |
| ``` | |
| ### 步骤3:重新训练 | |
| ```bash | |
| 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. 增加前景权重 | |
| ```json | |
| "foreground_weight": 5.0, | |
| "focal_alpha": [1.0, 5.0] | |
| ``` | |
| ### 2. 调整损失函数权重 | |
| ```json | |
| "dice_weight": 0.7, // 增加Dice Loss权重 | |
| "focal_weight": 1.0 | |
| ``` | |
| ### 3. 使用难例挖掘 | |
| - 重点关注模型误判的样本 | |
| - 增加这些样本在训练中的权重 | |
| ### 4. 检查数据质量 | |
| - 确保标签正确 | |
| - 检查是否有误标注 | |
| - 确保负样本(无浒苔图片)的标签是全0 | |
| ## 验证方法 | |
| 训练后,检查: | |
| 1. **训练损失**:应该逐渐下降 | |
| 2. **验证IoU**:应该逐渐上升 | |
| 3. **预测结果**:无浒苔图片应该被正确预测为背景 | |
| 4. **混淆矩阵**:假阳性(FP)应该减少 | |
| ## 注意事项 | |
| 1. **权重不要过大**:过大的权重可能导致训练不稳定 | |
| 2. **监控训练过程**:如果损失不下降,可能需要调整权重 | |
| 3. **验证集评估**:使用验证集来评估改进效果 | |