能否在Keras中实现基于批次类别不平衡的动态类别权重?
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_nanor 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.functionto 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

