File size: 3,660 Bytes
c2b1b26 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 | # 训练改进说明
## 问题诊断
即使有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. **验证集评估**:使用验证集来评估改进效果
|