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

如何构建分层采样的TensorFlow数据集以处理不平衡分类问题

结论

你当前的代码生成的批次不满足分层采样要求。全局shuffle加普通batch的逻辑只能实现样本随机打乱,每个批次内的类别比例会随机波动,大概率和训练集的不平衡原始分布一致,无法做到可控的分层采样。

调整方案

要实现按设定比例的分层采样批次,可以通过TensorFlow的tf.data.Dataset.sample_from_datasets接口实现,操作步骤如下:

  1. 先将训练集按类别拆分为独立的数据集(以下以二分类场景为例,多分类按相同逻辑扩展即可)
# 提取不同类别的样本索引
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()
  1. 按你的分层需求设置采样权重:
  • 若需要每个批次类别平衡:设置权重为[0.5, 0.5],多分类场景下每个类的权重为1/类别总数
  • 若需要每个批次的类别比例和训练集原始分布完全一致:设置权重为各类别样本占训练集总样本的比例
  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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 05:36:04