使用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
相关产品推荐
相关产品推荐

