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

能否在Keras中实现基于批次类别不平衡的动态类别权重?

Dynamic Class Weights Per Batch in Keras: Is It Feasible?

Absolutely yes! This approach is not only feasible but can also outperform static class_weights in scenarios where class distributions shift across batches—like streaming data, or imbalanced datasets where local batch ratios differ sharply from the global dataset. Here’s how to implement it in practice:

Core Idea

Instead of precomputing weights using the entire dataset’s class distribution, you calculate weights on-the-fly for each batch using the actual class counts present in that batch. These weights are then applied to the loss function, giving more importance to underrepresented classes in the current batch.

Implementation Methods

1. Custom Weighted Loss Function

The simplest way to integrate dynamic weights is to embed the calculation directly into your loss function. This works seamlessly with Keras’ built-in training loops.

For multi-class classification with one-hot encoded labels:

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

def dynamic_weighted_cce(y_true, y_pred):
    # Count samples per class in the current batch
    class_counts = tf.reduce_sum(y_true, axis=0)
    # Calculate weights: total batch size / (number of classes * class count)
    # Use divide_no_nan to avoid division by zero if a class is missing from the batch
    class_weights = tf.math.divide_no_nan(
        tf.cast(tf.shape(y_true)[0], tf.float32),
        tf.cast(class_counts * tf.shape(y_true)[1], tf.float32)
    )
    # Assign each sample the weight of its true class
    sample_weights = tf.reduce_sum(y_true * class_weights, axis=1)
    # Compute unweighted cross-entropy (reduction=NONE to get per-sample loss)
    cce = CategoricalCrossentropy(from_logits=False, reduction=tf.keras.losses.Reduction.NONE)
    per_sample_loss = cce(y_true, y_pred)
    # Return weighted average loss
    return tf.reduce_mean(per_sample_loss * sample_weights)

Compile your model with this loss function:

model.compile(optimizer='adam', loss=dynamic_weighted_cce)

If you’re using integer-encoded labels (shape (batch_size,)), adjust the class count calculation with tf.bincount:

def dynamic_weighted_sparse_cce(y_true, y_pred, num_classes):
    class_counts = tf.bincount(y_true, minlength=num_classes)
    class_weights = tf.math.divide_no_nan(
        tf.cast(tf.shape(y_true)[0], tf.float32),
        tf.cast(class_counts * num_classes, tf.float32)
    )
    # Map each sample's label to its corresponding weight
    sample_weights = tf.gather(class_weights, y_true)
    sparse_cce = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=False, reduction=tf.keras.losses.Reduction.NONE)
    per_sample_loss = sparse_cce(y_true, y_pred)
    return tf.reduce_mean(per_sample_loss * sample_weights)

2. Custom Training Loop (Full Control)

For more flexibility—like logging weight values or combining with custom preprocessing—use a custom training loop. This lets you explicitly compute weights before each batch step.

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

# Initialize components
optimizer = tf.keras.optimizers.Adam()
loss_fn = CategoricalCrossentropy(from_logits=False, reduction=tf.keras.losses.Reduction.NONE)
num_classes = 5  # Adjust to your dataset

@tf.function  # Optimize the training step with TensorFlow graph
def train_step(x_batch, y_batch):
    # Calculate dynamic class weights for the current batch
    class_counts = tf.reduce_sum(y_batch, axis=0)
    class_weights = tf.math.divide_no_nan(
        tf.cast(tf.shape(y_batch)[0], tf.float32),
        tf.cast(class_counts * num_classes, tf.float32)
    )
    sample_weights = tf.reduce_sum(y_batch * class_weights, axis=1)
    
    # Forward pass + gradient calculation
    with tf.GradientTape() as tape:
        y_pred = model(x_batch, training=True)
        per_sample_loss = loss_fn(y_batch, y_pred)
        weighted_loss = tf.reduce_mean(per_sample_loss * sample_weights)
    
    # Update model weights
    gradients = tape.gradient(weighted_loss, model.trainable_variables)
    optimizer.apply_gradients(zip(gradients, model.trainable_variables))
    
    return weighted_loss

# Run training
epochs = 10
for epoch in range(epochs):
    print(f"Epoch {epoch+1}/{epochs}")
    total_loss = 0.0
    batch_count = 0
    for x_batch, y_batch in train_dataset:
        loss = train_step(x_batch, y_batch)
        total_loss += loss.numpy()
        batch_count += 1
    print(f"Average Loss: {total_loss / batch_count:.4f}")

Key Considerations

  • Avoid Division by Zero: Always use tf.math.divide_no_nan or add a small epsilon (e.g., 1e-6) to class counts to prevent NaN values if a class is missing from a batch.
  • Weight Calculation Logic: The examples use the standard inverse frequency formula, but you can adjust this (e.g., 1 / (class_counts + 1) for smoother weights) based on your specific use case.
  • Performance: Dynamic weight calculation adds minimal overhead, especially when using tf.function to compile the training step into an optimized graph.

This approach ensures your model adapts to the class distribution of each batch, which can lead to better generalization in cases where the global dataset distribution doesn’t reflect local batch-level imbalances.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 08:45:00