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

Keras批次概率分布咨询:不平衡数据集CNN训练的批次分布维持方法

Balancing Batch Distributions for Imbalanced CNN Training in Keras

Great question—working with imbalanced datasets is a huge pain point for CNN training, and making sure your batches mirror the full training set's distribution is critical for avoiding unstable gradients and biased model performance. Let's break this down clearly:

1. How Well Do Batches Approximate the Training Set Distribution?

  • Default Sampling (Shuffled, No Replacement): Keras' default behavior (when you set shuffle=True in model.fit()) shuffles your entire training dataset each epoch, then splits it into sequential batches. For most cases, larger batch sizes will lead to batches that are closer to your full training set's class distribution. For example, if your data is 90% class A and 10% class B, a batch size of 256 will almost always have ~230 A samples and ~26 B samples, while a batch size of 32 might occasionally have only 1-2 B samples (or even none, in extreme imbalance cases).
  • Extreme Imbalance Edge Case: If your minority class makes up less than 5% of your data, even large batches can sometimes miss minority samples entirely—this is when you need intentional sampling strategies instead of relying on random shuffling.

2. Keras Tools to Match Batch Distribution to Training Set

Keras doesn't have a single "magic parameter" to enforce this, but there are straightforward built-in and TensorFlow-integrated ways to achieve it:

Option 1: Stratified Sampling with tf.data.Dataset

This is the most reliable method for controlling batch distribution. You can split your data by class, then sample from each class proportionally to its representation in the full training set:

import tensorflow as tf

# Assume you have x_train (input data) and y_train (class labels, integer-encoded)
class_counts = tf.math.bincount(y_train)
total_samples = len(y_train)

# Split data into separate datasets for each class
class_datasets = []
for class_idx in range(len(class_counts)):
    class_mask = tf.equal(y_train, class_idx)
    class_data = tf.data.Dataset.from_tensor_slices(
        (x_train[class_mask], y_train[class_mask])
    )
    class_datasets.append(class_data)

# Calculate sampling weights matching training set proportions
sampling_weights = [count / total_samples for count in class_counts]

# Create a dataset that samples proportionally from each class
balanced_dataset = tf.data.experimental.sample_from_datasets(
    class_datasets, weights=sampling_weights, seed=42
)

# Batch and prefetch for efficient training
balanced_dataset = balanced_dataset.batch(32).prefetch(tf.data.AUTOTUNE)

# Train your model with this balanced dataset
model.fit(balanced_dataset, epochs=15)

Option 2: Custom Sequence Generator (For Non-TensorFlow Data Sources)

If you're using a custom data pipeline (e.g., loading images from disk with Sequence), you can implement stratified sampling directly in the __getitem__ method:

from tensorflow.keras.utils import Sequence
import numpy as np

class StratifiedSequence(Sequence):
    def __init__(self, x, y, batch_size):
        self.x = x
        self.y = y
        self.batch_size = batch_size
        # Group indices by class
        self.class_indices = {}
        for class_idx in np.unique(y):
            self.class_indices[class_idx] = np.where(y == class_idx)[0]
        # Calculate number of samples to take per class per batch
        self.class_batch_sizes = {
            c: int(np.round(batch_size * (len(inds)/len(y))))
            for c, inds in self.class_indices.items()
        }

    def __len__(self):
        return len(self.x) // self.batch_size

    def __getitem__(self, idx):
        # Sample indices from each class proportionally
        batch_indices = []
        for c, batch_size in self.class_batch_sizes.items():
            batch_indices.extend(
                np.random.choice(self.class_indices[c], batch_size, replace=False)
            )
        # Shuffle to mix classes in the batch
        np.random.shuffle(batch_indices)
        return self.x[batch_indices], self.y[batch_indices]

# Use this generator in training
stratified_gen = StratifiedSequence(x_train, y_train, batch_size=32)
model.fit(stratified_gen, epochs=15)

Option 3: Class Weighting (Bonus: Loss-Level Balance)

While this doesn't change batch distribution, class_weight in model.fit() is a complementary tool. It assigns higher loss weights to minority classes, making the model prioritize learning from them even if batches are imbalanced:

# Calculate class weights (inverse of class frequency)
class_weights = {
    c: total_samples / (len(self.class_indices[c]) * len(self.class_indices))
    for c in self.class_indices
}

model.fit(x_train, y_train, batch_size=32, shuffle=True, class_weight=class_weights, epochs=15)

Key Takeaways

  • Batch distribution similarity depends on batch size and sampling strategy—larger batches are more reliable, but extreme imbalance needs intentional sampling.
  • Use tf.data.experimental.sample_from_datasets or a custom Sequence to enforce batch distributions that match your training set.
  • Pair data-level balancing with class_weight for extra robustness in imbalanced scenarios.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.20 08:06:22