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

如何通过回调更新ImageDataGenerator路径解决Keras多输入模型数据集不均衡问题

可行性结论

你提出的方案完全可以实现,这种动态欠采样轮换的思路既解决了单轮训练的类别不均衡问题,又能全量利用多数类的所有样本信息,不会浪费有效数据。

实现方式推荐

不要额外通过Callback调用ImageDataGenerator的方法,更简洁的方案是直接继承tf.keras.utils.Sequence类自定义生成器,Sequence类本身内置on_epoch_end方法,每轮训练结束后会自动触发执行,不需要额外的回调逻辑。
核心实现逻辑如下:

  • 初始化生成器时传入全量数据集DataFrame,统计三个类别的样本数量,取最少类的样本量作为每轮每个类的采样上限。
  • 在on_epoch_end方法中,对两个多数类分别做无放回随机采样,采样数量等于最少类的样本量,和少数类全量样本拼接后打乱,作为当前epoch的训练数据源。
  • 多轮训练后多数类的所有样本都会被轮询用到,同时每轮训练的类别分布都是均衡的。
  • 如果是多输入模型,只需要调整__getitem__方法的返回值,返回多个输入对应的数组列表即可,整体逻辑不需要修改。

如果你希望尽量复用flow_from_dataframe的能力,也可以在自定义Callback的on_epoch_end方法中生成新的均衡采样DataFrame,重新实例化flow_from_dataframe对象替换训练用的生成器,但这种方式性能不如自定义Sequence高效。

官方现成方案说明

目前TensorFlow/Keras没有完全匹配该需求的封装好的现成API,不过有两个替代方案可以参考:

  • 使用tf.data.Dataset API实现:将三个类别的数据分别封装为独立的Dataset对象,对多数类Dataset做shuffle和repeat操作,通过sample_from_datasets方法设置三类采样权重为[1/3,1/3,1/3],即可得到均衡采样的训练数据集,也能覆盖全量样本。
  • 训练时直接传入class_weight参数给少数类设置更高的损失权重,不过这种方案在类别不均衡差距较大的情况下效果不如动态采样稳定。
简化代码示例
import pandas as pd
import numpy as np
from tensorflow.keras.utils import Sequence
from tensorflow.keras.preprocessing.image import load_img, img_to_array

class BalancedMultiInputGenerator(Sequence):
    def __init__(self, total_df, img_col='path', label_col='label', 
                 extra_input_cols=['feat1', 'feat2'], # 多输入的其他特征列
                 batch_size=32, target_size=(224,224), aug=None):
        self.total_df = total_df
        self.img_col = img_col
        self.label_col = label_col
        self.extra_input_cols = extra_input_cols
        self.batch_size = batch_size
        self.target_size = target_size
        self.aug = aug
        # 统计最小类样本量作为每轮采样上限
        self.class_counts = self.total_df[self.label_col].value_counts()
        self.min_sample_num = self.class_counts.min()
        self.class_list = self.class_counts.index.tolist()
        # 初始化第一轮数据集
        self.on_epoch_end()
    
    def __len__(self):
        # 返回每轮的batch总数
        return int(np.ceil(len(self.current_epoch_df) / self.batch_size))
    
    def __getitem__(self, idx):
        # 取当前batch的数据
        batch_df = self.current_epoch_df.iloc[idx*self.batch_size : (idx+1)*self.batch_size]
        img_input = []
        extra_input = []
        labels = []
        for _, row in batch_df.iterrows():
            # 处理图像输入
            img = load_img(row[self.img_col], target_size=self.target_size)
            img = img_to_array(img) / 255.0
            if self.aug is not None:
                img = self.aug.random_transform(img)
            img_input.append(img)
            # 处理其他输入
            extra_input.append(row[self.extra_input_cols].values.astype(np.float32))
            labels.append(row[self.label_col])
        # 多输入返回列表,单输入可直接返回第一个数组
        return [np.array(img_input), np.array(extra_input)], np.array(labels)
    
    def on_epoch_end(self):
        # 每轮结束重新采样均衡数据集
        sampled_dfs = []
        for cls in self.class_list:
            cls_all_df = self.total_df[self.total_df[self.label_col] == cls]
            # 无放回采样,确保每轮样本不重复,多轮覆盖全量
            sampled_cls_df = cls_all_df.sample(n=self.min_sample_num, replace=False)
            sampled_dfs.append(sampled_cls_df)
        # 合并后打乱顺序
        self.current_epoch_df = pd.concat(sampled_dfs).sample(frac=1).reset_index(drop=True)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 22:45:03