如何构建分层采样的TensorFlow数据集以处理不平衡分类问题
结论
你当前的代码生成的批次不满足分层采样要求。全局shuffle加普通batch的逻辑只能实现样本随机打乱,每个批次内的类别比例会随机波动,大概率和训练集的不平衡原始分布一致,无法做到可控的分层采样。
调整方案
要实现按设定比例的分层采样批次,可以通过TensorFlow的tf.data.Dataset.sample_from_datasets接口实现,操作步骤如下:
- 先将训练集按类别拆分为独立的数据集(以下以二分类场景为例,多分类按相同逻辑扩展即可)
# 提取不同类别的样本索引 class_0_idx = y_train == 0 class_1_idx = y_train == 1 # 分别构造每个类别的数据集,独立shuffle后设置repeat避免样本少的类别提前跑完 ds_class0 = tf.data.Dataset.from_tensor_slices((X_train[class_0_idx], y_train[class_0_idx])).shuffle(buffer_size=sum(class_0_idx)).repeat() ds_class1 = tf.data.Dataset.from_tensor_slices((X_train[class_1_idx], y_train[class_1_idx])).shuffle(buffer_size=sum(class_1_idx)).repeat()
- 按你的分层需求设置采样权重:
- 若需要每个批次类别平衡:设置权重为
[0.5, 0.5],多分类场景下每个类的权重为1/类别总数 - 若需要每个批次的类别比例和训练集原始分布完全一致:设置权重为各类别样本占训练集总样本的比例
- 合并数据集并打批次
# 示例为平衡采样的权重,可按需调整 sample_weights = [0.5, 0.5] dataset = tf.data.Dataset.sample_from_datasets( [ds_class0, ds_class1], weights=sample_weights ).batch(batch_size)
注意事项
- 上述代码中给每个类别数据集加了
repeat(),会让数据集无限迭代,训练时需要手动控制每轮训练的步数,设置为X_train.shape[0] // batch_size即可和普通训练的单epoch长度保持一致 - 你原有代码中
train_test_split指定stratify=y_new的逻辑是正确的,已经保证了训练集、测试集的类别比例和原数据集一致,无需调整
内容的提问来源于stack exchange,提问作者RasM10
相关产品推荐
相关产品推荐

