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` 来:
- 统计类别分布
- 计算建议的类别权重
- 检查数据平衡情况
```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. **验证集评估**:使用验证集来评估改进效果