发布于2026-07-19 阅读(0)
扫一扫,手机访问
在深度学习中,类别不平衡是个老生常谈的问题。加权损失函数是最直接有效的应对手段之一,但很多人在TensorFlow里踩过坑——不是效果不明显,就是验证集指标突然崩了。今天就把这些坑填平,说说加权损失函数到底该怎么用。

tf.keras.losses.BinaryCrossentropy 加权处理二分类不平衡不少人第一反应是往损失函数初始化里塞 class_weight 参数,但这条路走不通——class_weight 是 model.fit() 的专属参数,跟损失函数本身没关系。真正能让损失函数内部对正负样本分别加权的,得手动改计算逻辑,或者换用更底层的封装方式。
最稳妥的做法是自定义损失函数,有两个主流方向:
tf.nn.weighted_cross_entropy_with_logits,但它只接受未经过 sigmoid 的 logits 输入,且权重只作用于正样本(label=1),负样本默认权重固定为 1。如果想两边都控制权重,就得自己实现。tf.keras.losses.binary_crossentropy 算出每个样本的 loss 后,按真实 label 选择对应权重相乘,最后取均值。注意 label 必须显式转为 float32,否则和权重相乘时可能触发隐式类型转换错误。tf.keras.losses.CategoricalCrossentropy 加权多分类加权不能指望靠缩放输出层搞定,必须在损失计算中对每个类别的样本赋予不同权重。TensorFlow 原生不支持直接给损失函数传入 per-class 权重,但可以通过 sample_weight 机制间接实现。
核心操作:训练时把每个样本的单维权重向量(长度 = batch_size)传给 model.fit(..., sample_weight=...)。这个权重向量需要根据每个样本的真实类别索引查表得到:
[1/0.9, 1/0.1],或者直接用 sklearn 的 compute_class_weight 自动生成。sample_weight 数组:对 batch 中每个样本,用其 one-hot label 的 tf.argmax 找到类别 id,再查权重数组取值。sample_weight 的 shape 是 (batch_size,),不能是 (batch_size, num_classes),否则会抛出 ValueError: sample_weight tensor must be 1D。class_weight 参数有时没效果?class_weight 本质上是 model.fit() 的高级封装,底层还是靠生成 sample_weight 实现。它失效通常有以下几个原因:
activation='sigmoid' 但标签是整数(如 [0, 1, 1, 0]),此时 Keras 默认当作 multi-label 处理,class_weight 不生效——必须转成 one-hot 或改用 sparse 版本损失。SparseCategoricalCrossentropy 损失,但 class_weight 传的是 dict 形式(如 {0: 1.0, 1: 5.0}),这其实有效;但如果传成了 list 就会静默失败。tf.data.Dataset 且没把权重作为第三项返回(即 dataset.map(lambda x, y: (x, y, weight))),class_weight 完全不起作用。加权只影响训练时梯度更新方向,不影响验证阶段 loss 计算逻辑。如果发现 val_loss 突然变大、val_accuracy 剧烈波动,大概率是验证集也误加了权重。
务必确认:model.evaluate() 和 model.fit(validation_data=...) 中的验证数据不能带 sample_weight,否则评估结果失真:
tf.data.Dataset 构建验证集,确保 map 函数只返回 (x, y),不要返回第三个权重项。class_weight 时,Keras 默认不会对 validation_data 应用权重——但如果你手动传了 validation_sample_weight,就得自己负责清零。类别不平衡问题里,最容易被忽略的不是怎么加权,而是加权后如何评估。auc、f1-score、per-class recall 这些指标比 accuracy 更能反映模型在少数类上的真实能力。别忘了。
售后无忧
立即购买>office旗舰店
售后无忧
立即购买>office旗舰店
售后无忧
立即购买>office旗舰店
售后无忧
立即购买>office旗舰店
正版软件
正版软件
正版软件
正版软件
正版软件
1
2
3
7
8