Keras中如何处理语义分割任务的类别不平衡?附类FCN-32s模型代码
针对你用Keras搭建的基于预训练VGG16的FCN-32s模型,处理语义分割任务中的类别不平衡问题,我整理了几个在Keras生态里容易落地的实用方案,按优先级和易用性排序:
1. 直接使用类别权重(Class Weights)
这是最省心的入门方案,Keras的fit()/fit_generator()方法原生支持class_weight参数。你只需要先统计训练数据中每个类别的像素占比(语义分割里按像素统计比按样本数更准确),然后计算每个类别的权重,公式一般是:权重 = 总像素数 / (类别数量 * 该类别像素数)
举个具体的实现例子:
import numpy as np from sklearn.utils.class_weight import compute_class_weight # 假设你的标签是整数形式(0,1,2,...,n_classes-1) y_flatten = y_train.flatten() class_weights = compute_class_weight('balanced', classes=np.unique(y_flatten), y=y_flatten) class_weight_dict = dict(zip(np.unique(y_flatten), class_weights)) # 训练时传入权重字典 model.fit(x_train, y_train, class_weight=class_weight_dict, epochs=20, batch_size=8)
这个方法适合类别不平衡程度中等的场景,不需要修改模型或损失函数,成本极低。
2. 自定义加权交叉熵损失(Weighted Cross-Entropy)
如果类别不平衡非常极端(比如某类像素占比不到1%),类别权重的效果可能不够,这时候可以给交叉熵损失加上类别权重。在Keras里可以轻松自定义这个损失函数:
import tensorflow.keras.backend as K def weighted_categorical_crossentropy(weights): def loss(y_true, y_pred): # y_true是one-hot编码的标签,shape=(batch, height, width, n_classes) cross_ent = K.categorical_crossentropy(y_true, y_pred) # 提取每个像素对应的类别权重 weight = K.sum(weights * y_true, axis=-1) # 返回加权后的损失 return cross_ent * weight return loss # 假设weights是长度为n_classes的数组,对应每个类别的权重 model.compile(optimizer='adam', loss=weighted_categorical_crossentropy(class_weights))
如果你的标签是稀疏整数形式(不是one-hot),可以把K.categorical_crossentropy换成K.sparse_categorical_crossentropy,并调整权重的提取逻辑。
3. 数据层面的采样策略
从数据入手缓解不平衡,也是常用的思路:
- 过采样(Oversampling):对小类别的样本/像素进行数据增强(翻转、旋转、缩放、亮度调整等),或者直接复制小类别占比高的图像。你可以在Keras的
ImageDataGenerator里针对小类别样本单独增强,或者写自定义数据生成器,在生成batch时优先采样小类别样本。 - 欠采样(Undersampling):对大类别的样本随机丢弃一部分,这种方法简单但会浪费数据,只适合大类样本量远超小类的极端场景,不推荐作为首选。
- 混合采样:结合两者,比如对小类别过采样2倍,对大类别欠采样到原来的50%,平衡各类别的数据量。
4. 使用Focal Loss
Focal Loss专门针对类别不平衡和难分类样本设计,它通过降低易分类样本(大类中的多数像素)的损失权重,让模型更关注难分类的小类别样本。Keras里的实现示例:
import tensorflow as tf import tensorflow.keras.backend as K def focal_loss(gamma=2.0, alpha=None): def focal_loss_fixed(y_true, y_pred): # 确保y_pred在(0,1)之间,避免log(0)报错 y_pred = K.clip(y_pred, K.epsilon(), 1 - K.epsilon()) # 计算交叉熵 cross_ent = -y_true * K.log(y_pred) # 计算调制因子,降低易分类样本的权重 pt = tf.where(tf.equal(y_true, 1), y_pred, 1 - y_pred) loss = alpha * K.pow(1 - pt, gamma) * cross_ent return K.sum(loss, axis=-1) return focal_loss_fixed # 多分类场景下,alpha可以是对应每个类别的权重数组;二分类可设为0.25 model.compile(optimizer='adam', loss=focal_loss(gamma=2, alpha=class_weights))
Focal Loss在极端不平衡的场景下效果通常比加权交叉熵更好,但需要调整gamma和alpha两个超参数。
5. 选择合适的评价指标
别只看准确率(Accuracy)!在类别不平衡的情况下,准确率会被大类主导,完全不能反映模型对小类的性能。你应该使用IoU(交并比)、Dice系数、F1-Score这些更鲁棒的指标,在Keras里可以自定义这些指标:
def dice_coef(y_true, y_pred): smooth = 1e-6 y_true_f = K.flatten(y_true) y_pred_f = K.flatten(y_pred) intersection = K.sum(y_true_f * y_pred_f) return (2. * intersection + smooth) / (K.sum(y_true_f) + K.sum(y_pred_f) + smooth) def iou(y_true, y_pred): smooth = 1e-6 intersection = K.sum(y_true * y_pred) union = K.sum(y_true) + K.sum(y_pred) - intersection return (intersection + smooth) / (union + smooth) # 编译模型时加入这些指标 model.compile(optimizer='adam', loss=..., metrics=[dice_coef, iou])
总结一下实践顺序
建议先从类别权重开始尝试,成本最低;如果效果不好,再换加权交叉熵;极端不平衡场景下试试Focal Loss;同时配合数据增强/采样和合适的评价指标,基本能覆盖大部分语义分割的类别不平衡问题。
内容的提问来源于stack exchange,提问作者mrgloom

