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

Keras中fit_generator如何为每个训练批次设置class_weight?

How to Apply Per-Batch Class Weights in Keras fit_generator for Segmentation Tasks

Great question—this is a common pain point when dealing with imbalanced segmentation datasets, and the key here is to shift from using global class_weight to generating per-sample weights for each batch directly in your generator. Here's why your initial approach didn't work, and how to fix it without reinventing the wheel:

Why Your Generator Approach Threw an Error

When you tried returning class_weight from your generator, Keras got confused because fit_generator expects generators to yield one of these formats:

  • (x_batch, y_batch)
  • (x_batch, y_batch, sample_weights_batch)

A class_weight dictionary isn't a valid return value here. The TypeError: object of type 'generator' has no len() likely popped up because Keras tried to infer the generator's length (for progress tracking) and couldn't, possibly due to the invalid return format breaking internal checks.

The Correct Approach: Per-Batch Sample Weights

For segmentation tasks, since we're dealing with per-pixel labels, we need to create a sample weight array that matches the shape of our segmentation mask. Each pixel's weight is determined by the class frequency in the current batch. Here's how to implement this in your generator:

Step 1: Modify Your Generator to Compute Per-Batch Weights

Inside your generator, after generating a batch of inputs (x_batch) and segmentation masks (y_batch), calculate the class frequencies for that batch, then create a weight array where each pixel gets the weight of its class.

Example for Integer Label Masks (shape: (batch_size, height, width))

import numpy as np

def your_segmentation_generator(...):
    while True:
        # Generate your batch of x and y (integer labels)
        x_batch, y_batch = ...  # Your existing code to load/preprocess data
        
        # Calculate class frequencies in the current batch
        flat_mask = y_batch.flatten()
        class_counts = np.bincount(flat_mask)
        
        # Avoid division by zero for classes not present in the batch
        class_counts[class_counts == 0] = 1
        
        # Compute inverse frequency weights (adjust this logic to your needs)
        class_weights = 1.0 / class_counts
        
        # Create sample weight array matching the mask shape
        sample_weights = class_weights[flat_mask].reshape(y_batch.shape)
        
        # Yield the batch with sample weights
        yield (x_batch, y_batch, sample_weights)

Example for One-Hot Encoded Masks (shape: (batch_size, height, width, num_classes))

If your masks are one-hot encoded, first convert them to integer labels to compute weights:

def your_segmentation_generator(...):
    while True:
        x_batch, y_batch_onehot = ...  # Your existing code
        
        # Convert one-hot to integer labels
        y_batch = np.argmax(y_batch_onehot, axis=-1)
        
        # Same weight calculation as above
        flat_mask = y_batch.flatten()
        class_counts = np.bincount(flat_mask)
        class_counts[class_counts == 0] = 1
        class_weights = 1.0 / class_counts
        sample_weights = class_weights[flat_mask].reshape(y_batch.shape)
        
        # If your model expects one-hot labels, yield the original one-hot mask
        yield (x_batch, y_batch_onehot, sample_weights)

Step 2: Train with fit_generator (No class_weight Needed)

Now, when you call fit_generator, you don't need to pass the class_weight parameter—your generator is already providing per-sample weights tailored to each batch:

model.fit_generator(
    generator=your_segmentation_generator(...),
    steps_per_epoch=...,
    epochs=...,
    # No class_weight argument needed here
)

What About Using Callbacks?

You mentioned considering custom callbacks, but this approach is unnecessary here because modifying the generator is simpler and more direct. Callbacks in Keras don't receive the raw batch data by default (the logs dictionary only contains metrics, not inputs/labels), so you'd have to hack around storing batch data globally or using a custom training loop—something that's more work than just adjusting your generator.

Adjusting the Weight Calculation

The example uses inverse class frequency, but you can tweak this to match your needs:

  • Use normalized weights (divide by the sum of weights to keep the total loss scale consistent)
  • Use a more sophisticated formula like class_weights = total_pixels / (num_classes * class_counts)
  • Add a smoothing factor to avoid extreme weights for rare classes (e.g., class_counts = class_counts + 1e-6)

This method ensures each batch uses weights based on its own class distribution, which is exactly what you're looking for—no reinventing the wheel, just leveraging Keras's built-in support for sample weights in generators.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 06:24:47