加权交叉熵比直接重采样更可控,因其将惩罚力度嵌入损失函数,不改变数据分布或引入噪声;权重须基于训练集真实分布计算,且需通过sample_weight而非class_weight传入。

为什么加权交叉熵比直接重采样更可控
类别不平衡时,模型容易偏向多数类,单纯用 class_weight='balanced' 或过采样少数类,常导致验证集准确率虚高、F1暴跌。加权交叉熵把“惩罚力度”直接嵌入损失函数,训练时对每个样本的梯度更新按类别重要性缩放,不改变数据分布,也不引入合成样本噪声。
关键点在于:权重不是凭经验拍的,得从训练集真实分布算——比如二分类中,若正样本占比 5%,则正类权重应设为 len(y_train) / (2 * len(y_train[y_train==1])),分母的 2 是类别数,避免权重过大导致 loss 爆炸。
- 权重必须在训练前计算,不能用验证集统计值
- TensorFlow 2.x 中推荐用
tf.keras.losses.BinaryCrossentropy(from_logits=True)+ 手动加权,而非旧版weighted_cross_entropy_with_logits(已弃用) - 若用
from_logits=False,需确保输出已过 sigmoid,否则权重会和激活函数相互干扰
如何在 tf.keras.Model.fit() 中正确传入样本权重
很多人误以为设置 class_weight 参数就够了,但 class_weight 只支持整数标签且仅作用于 fit 的顶层逻辑;若模型含自定义层、多输出或使用 tf.data.Dataset,必须显式构造 sample_weight 数组并传入。
实操步骤:
立即学习“Python免费学习笔记(深入)”;
- 先用
np.bincount(y_train)得到各类频次,再按公式weight = total_samples / (n_classes * class_count)算出每个类的权重值 - 用
np.array([class_weights[y] for y in y_train])构造与y_train同长的sample_weight数组 - 调用
model.fit(X_train, y_train, sample_weight=sample_weight, ...)——注意不是class_weight - 若用
tf.data.Dataset.from_tensor_slices,需用zip((X, y), sample_weight)包装三元组
加权交叉熵的两种实现方式及性能差异
TensorFlow 提供两种加权路径:一是用内置 tf.keras.losses.BinaryCrossentropy 的 sample_weight 参数,二是手写带权重的 loss 函数。前者简洁稳定,后者灵活但易出错。
推荐用内置方式:
loss_fn = tf.keras.losses.BinaryCrossentropy(from_logits=True) # 在 model.compile 中指定 model.compile(optimizer='adam', loss=loss_fn, metrics=['accuracy']) # fit 时传 sample_weight,loss_fn 自动应用权重
手写 loss 示例(仅当需动态权重时考虑):
def weighted_bce(y_true, y_pred):
weights = tf.where(y_true == 1, pos_weight, 1.0)
unweighted_loss = tf.keras.losses.binary_crossentropy(y_true, y_pred, from_logits=True)
return tf.reduce_mean(weights * unweighted_loss)
注意:pos_weight 必须是标量 tf.constant 或张量,不能是 Python float,否则会触发 Eager 模式下重复图构建,拖慢训练。
验证阶段指标失真问题怎么绕过去
加权 loss 只影响训练梯度,不影响验证指标计算逻辑——val_loss 默认仍用未加权 loss,而 val_accuracy 对不平衡数据毫无意义。必须手动加入 F1、precision、recall 等指标,且验证集 sample_weight 应保持为全 1(即不加权),否则评估结果不可比。
- 用
tf.keras.metrics.F1Score(TF 2.13+)或自定义tf.keras.metrics.Metric子类 - 验证时不要传
sample_weight,哪怕训练用了——验证目标是看模型泛化能力,不是看加权拟合程度 - 如果用
class_weight而非sample_weight,验证指标默认不加权,这点反而更干净
最常被忽略的是:加权后模型输出 logits 偏移,直接用 0.5 截断阈值会导致召回率骤降,务必在验证集上用 precision_recall_curve 重选阈值。


















