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

使用imblearn的BalancedBatchGenerator返回类别数不一致引发训练报错

问题原因

imblearn.keras.BalancedBatchGenerator 采用随机采样逻辑,当batch_size未覆盖所有类别采样阈值时,可能出现单批次遗漏部分小类别样本的情况,自动生成的标签独热编码维度会随批次内实际类别数变化,和模型输出层固定的12维节点不匹配,触发维度广播错误。

解决方案
  • 方案1:调整batch大小为类别数的整数倍

    12分类任务中,将batch_size设置为12的整数倍(如36、48),保证每类在平衡采样逻辑下至少能分配到1个样本名额,从概率层面大幅降低漏类概率。
    修改生成器初始化代码:

    gen = BalancedBatchGenerator(data[["image"]], data[["class"]],
                                    sampler=RandomOverSampler(),
                                    random_state = 10,
                                    batch_size=36) # 12*3,每类平均3个样本
    
  • 方案2:手动固定独热编码维度(最稳妥)

    不依赖生成器自动生成独热标签,自行封装生成器强制指定标签维度为12,哪怕批次内缺少某类样本,标签维度也能和模型输出匹配:

    import tensorflow as tf
    
    def fixed_label_generator(original_gen, num_classes=12):
        for x, y in original_gen:
            # 手动转独热,固定维度为12
            y_fixed = tf.keras.utils.to_categorical(y, num_classes=num_classes)
            yield x, y_fixed
    
    # 替换原生成器
    train_generator = fixed_label_generator(gen, num_classes=12)
    
  • 方案3:自定义强制全类采样的生成器

    完全替换原生BalancedBatchGenerator,自行实现分层采样逻辑,每次生成批次时强制从每个类别采样相同数量的样本,保证每个批次都包含全部12个类别:

    import random
    import numpy as np
    import pandas as pd
    
    # 提前按类别分组
    class_groups = data.groupby("class")
    class_list = list(class_groups.groups.keys())
    per_class_sample = 3 # 每类采样3个,总batch_size=12*3=36
    def custom_balanced_gen():
        while True:
            batch_x = []
            batch_y = []
            # 每类抽固定数量样本
            for cls in class_list:
                cls_samples = class_groups.get_group(cls).sample(n=per_class_sample, replace=True)
                batch_x.extend(cls_samples["image"].tolist())
                batch_y.extend(cls_samples["class"].tolist())
            # 打乱批次内样本顺序
            combined = list(zip(batch_x, batch_y))
            random.shuffle(combined)
            batch_x, batch_y = zip(*combined)
            # 转成模型输入格式并生成固定维度独热标签
            yield np.array(batch_x), tf.keras.utils.to_categorical(batch_y, num_classes=12)
    
  • 方案4:改用类权重替代平衡生成器(改动最小)

    如果不需要强制batch内类别平衡,可直接去掉BalancedBatchGenerator,改用普通批次生成器配合类权重参数解决类别不平衡问题,完全避免漏类维度错误:

    from sklearn.utils.class_weight import compute_class_weight
    
    # 计算平衡类权重
    all_labels = data["class"].values
    class_weights = compute_class_weight("balanced", classes=np.unique(all_labels), y=all_labels)
    class_weight_dict = {idx: weight for idx, weight in enumerate(class_weights)}
    
    # 训练时传入类权重
    history = model.fit(
        generator = train_generator, 
        validation_data = val_generator,
        epochs = 50,
        verbose = 1,
        callbacks = callbacks,
        class_weight = class_weight_dict
    )
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.07 04:42:00