You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

在Keras或TensorFlow中仅为误分类二分类样本加权的实现方法

Weighting Misclassified Samples in Binary Cross-Entropy (Keras/TensorFlow)

Absolutely, you can target and weight misclassified samples from a specific class in your binary classification task—you just can’t do it with the out-of-the-box class_weight or sample_weight alone. Those parameters apply static weights (set before training starts), but weighting only misclassified samples requires dynamic logic that reacts to the model’s predictions during training.

Here’s how to implement this with a custom loss function:

Core Approach

The idea is to modify the standard Binary Cross-Entropy (BCE) loss to:

  1. Calculate the base BCE loss for every sample
  2. Identify which samples from your target class are misclassified
  3. Multiply the loss of those misclassified samples by your desired extra weight, while leaving other samples’ loss unchanged

Custom Loss Function Implementation

Let’s say you want to weight misclassified positive class samples (true label = 1, model predicts < 0.5) by a factor of 3.0. Here’s a TensorFlow/Keras-compatible function:

import tensorflow as tf
from tensorflow.keras.losses import BinaryCrossentropy

def weighted_misclass_bce(target_class=1, misclass_weight=3.0):
    def loss(y_true, y_pred):
        # Calculate raw BCE loss without reduction (per-sample loss)
        base_bce = BinaryCrossentropy(reduction='none')(y_true, y_pred)
        
        # Define misclassification condition for the target class
        if target_class == 1:
            # True positives that were predicted as negative
            is_misclassified = tf.logical_and(
                tf.equal(y_true, 1.0),
                tf.less(y_pred, 0.5)
            )
        else:
            # True negatives that were predicted as positive
            is_misclassified = tf.logical_and(
                tf.equal(y_true, 0.0),
                tf.greater_equal(y_pred, 0.5)
            )
        
        # Create weight tensor: apply extra weight to misclassified samples, 1.0 otherwise
        weights = tf.where(is_misclassified, misclass_weight, 1.0)
        
        # Return the average of weighted per-sample losses
        return tf.reduce_mean(base_bce * weights)
    return loss

How to Use It

When compiling your model, pass the custom loss function (tweak target_class and misclass_weight to match your needs):

model.compile(
    optimizer='adam',
    loss=weighted_misclass_bce(target_class=1, misclass_weight=3.0),
    metrics=['accuracy']
)

Combining with Class Weights

If you also want to apply static class weights (e.g., to account for class imbalance) alongside the misclassification weighting, you can extend the loss function:

def weighted_misclass_bce_with_class_weights(
    class_weights={0: 1.0, 1: 2.0},
    target_class=1,
    misclass_weight=3.0
):
    def loss(y_true, y_pred):
        base_bce = BinaryCrossentropy(reduction='none')(y_true, y_pred)
        
        # Apply static class weights first
        class_weight_tensor = tf.where(
            tf.equal(y_true, 1.0),
            class_weights[1],
            class_weights[0]
        )
        
        # Add misclassification weighting
        if target_class == 1:
            is_misclassified = tf.logical_and(tf.equal(y_true, 1.0), tf.less(y_pred, 0.5))
        else:
            is_misclassified = tf.logical_and(tf.equal(y_true, 0.0), tf.greater_equal(y_pred, 0.5))
        
        misclass_weight_tensor = tf.where(is_misclassified, misclass_weight, 1.0)
        
        # Combine both weights
        total_weights = class_weight_tensor * misclass_weight_tensor
        
        return tf.reduce_mean(base_bce * total_weights)
    return loss

Key Notes

  • Threshold Adjustment: The example uses 0.5 as the classification threshold—if you’re using a different threshold (e.g., for precision/recall tradeoffs), update the tf.less/tf.greater_equal conditions accordingly.
  • Tensor Operations: All logic uses TensorFlow tensor operations (not Python conditionals inside the loss computation) to ensure compatibility with the training graph and automatic differentiation.
  • Dynamic Weighting: This approach weights misclassified samples per batch, meaning the model will adapt its focus as it learns (samples that were once misclassified but now get right will stop receiving extra weight).

内容的提问来源于stack exchange,提问作者Nickpick

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.19 07:33:15