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

Keras训练卷积神经网络时如何用flow_from_dataframe平衡小批量类别

实现方案

原生flow_from_dataframe没有直接提供按类别均衡采样的能力,你可以通过以下两种方案实现每个批量都包含三类样本的需求:

方案1:包装多生成器实现均衡采样

思路是按标签拆分数据集,每个类别单独创建生成器,每次批量训练时从三个生成器各取对应数量的样本拼接成完整批量。

步骤1:按标签拆分训练集

import numpy as np
import pandas as pd
from tensorflow.keras.preprocessing.image import ImageDataGenerator

# 按标签拆分为三个独立子集
train_0 = train[train['label'] == 0].reset_index(drop=True)
train_1 = train[train['label'] == 1].reset_index(drop=True)
train_2 = train[train['label'] == 2].reset_index(drop=True)

步骤2:为每个类别创建独立生成器

# 复用你原有的图像增强配置
datagen = ImageDataGenerator(rescale=1./255) # 保留你原本的增强参数即可

# 三个类别的生成器,批量大小加总为你需要的总批量64
gen0 = datagen.flow_from_dataframe(
    dataframe=train_0,
    directory=None,
    x_col="directory",
    y_col="label",
    batch_size=21,
    seed=42,
    shuffle=True,
    class_mode='categorical',
    target_size=(299,299)
)

gen1 = datagen.flow_from_dataframe(
    dataframe=train_1,
    directory=None,
    x_col="directory",
    y_col="label",
    batch_size=21,
    seed=42,
    shuffle=True,
    class_mode='categorical',
    target_size=(299,299)
)

gen2 = datagen.flow_from_dataframe(
    dataframe=train_2,
    directory=None,
    x_col="directory",
    y_col="label",
    batch_size=22,
    seed=42,
    shuffle=True,
    class_mode='categorical',
    target_size=(299,299)
)

步骤3:自定义组合生成器拼接批量

def balanced_batch_generator(gen0, gen1, gen2):
    while True:
        # 从每个生成器取对应数量的样本
        X0, y0 = next(gen0)
        X1, y1 = next(gen1)
        X2, y2 = next(gen2)
        # 拼接样本和标签
        batch_X = np.concatenate([X0, X1, X2], axis=0)
        batch_y = np.concatenate([y0, y1, y2], axis=0)
        # 打乱批量内样本顺序,避免同类样本扎堆
        shuffle_idx = np.random.permutation(len(batch_X))
        yield batch_X[shuffle_idx], batch_y[shuffle_idx]

# 初始化最终的训练生成器
train_generator = balanced_batch_generator(gen0, gen1, gen2)

训练注意事项

自定义生成器无法自动计算每轮迭代步数,需要手动指定steps_per_epoch,取值为三个生成器迭代步数的最小值,避免某类样本先跑完报错:

steps_per_epoch = min(len(gen0), len(gen1), len(gen2))
model.fit(
    train_generator,
    steps_per_epoch=steps_per_epoch,
    # 其余你原本的训练参数保持不变
)

方案2:使用tf.data API实现(更推荐)

TensorFlow 2.x的tf.data接口原生支持多数据集加权采样,性能更优、灵活性更高:

import tensorflow as tf

# 定义图像加载预处理逻辑,和你原本的预处理规则保持一致
def load_image(file_path, label):
    img = tf.io.read_file(file_path)
    img = tf.image.decode_jpeg(img, channels=3)
    img = tf.image.resize(img, (299, 299))
    img = img / 255.0
    label = tf.one_hot(label, depth=3)
    return img, label

# 为每个类别创建独立数据集
ds0 = tf.data.Dataset.from_tensor_slices((train_0['directory'].values, train_0['label'].values)).map(load_image, num_parallel_calls=tf.data.AUTOTUNE).shuffle(1000).repeat()
ds1 = tf.data.Dataset.from_tensor_slices((train_1['directory'].values, train_1['label'].values)).map(load_image, num_parallel_calls=tf.data.AUTOTUNE).shuffle(1000).repeat()
ds2 = tf.data.Dataset.from_tensor_slices((train_2['directory'].values, train_2['label'].values)).map(load_image, num_parallel_calls=tf.data.AUTOTUNE).shuffle(1000).repeat()

# 按权重采样,设置等权重即可保证每类都有样本被抽到
train_ds = tf.data.Dataset.sample_from_datasets(
    [ds0, ds1, ds2],
    weights=[1/3, 1/3, 1/3],
    seed=42
).batch(64).prefetch(tf.data.AUTOTUNE)

训练时直接传入train_ds即可,无需额外调整参数。

额外提示

如果你的类别不均衡程度极高,可通过调整每个类的采样权重、批量占比来平衡训练效果,同时注意少类样本不要过度重复采样避免过拟合。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 18:06:04