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

TensorFlow入门者咨询:不均衡数据集下采样实现方法

针对TensorFlow预定义Estimator处理极度不均衡数据集的建议

Hey there! 作为TensorFlow入门者碰到类别不均衡的问题太正常了,尤其是你这种正样本仅占0.1%的极端情况,先给你点个赞——能意识到数据不均衡会严重拖累模型性能,这已经迈出了关键一步!

既然你数据量充足,选择下采样负样本构建均衡数据集的思路非常稳妥(比上采样正样本更不容易触发过拟合),下面针对你提到的预处理方向,结合预定义Estimator的使用场景,给你两种具体实现方案和注意事项:

方案一:用tf.data.Dataset在数据加载阶段直接过滤采样

这是预定义Estimator官方推荐的数据流处理方式,适合数据量极大、无法全量放入内存的场景:

  • 核心思路:先把原始数据集拆分成正样本子集和负样本子集,然后从负样本子集中采样和正样本数量相等的样本,最后合并两个子集并打乱
  • 代码示例(假设你用CSV格式数据):
import tensorflow as tf

def create_balanced_dataset(file_path, batch_size=32):
    # 加载原始数据集,解析特征和标签
    dataset = tf.data.experimental.make_csv_dataset(
        file_path,
        batch_size=batch_size,
        label_name='label',  # 替换成你的标签列名
        num_epochs=1,
        shuffle=False  # 先不打乱,方便拆分正负样本
    )

    # 拆分正、负样本子集
    positive_ds = dataset.filter(lambda features, label: label == 1)
    negative_ds = dataset.filter(lambda features, label: label == 0)

    # 统计正样本数量(如果提前知道可以直接写死,省得遍历)
    pos_count = sum(1 for _ in positive_ds)
    # 采样负样本,数量和正样本一致
    sampled_negative_ds = negative_ds.take(pos_count)

    # 合并并打乱数据集
    balanced_ds = positive_ds.concatenate(sampled_negative_ds).shuffle(buffer_size=pos_count*2)
    return balanced_ds.repeat()  # 重复数据集供训练使用
  • 优点:完全基于TensorFlow的数据流API,和预定义Estimator兼容度拉满,不需要额外处理内存问题
  • 注意点:如果是TFRecord格式数据,只需要把make_csv_dataset换成对应的TFRecord加载逻辑即可;另外如果正样本数量动态变化,记得每次训练前重新统计数量。

方案二:在Estimator的输入函数(input_fn)中动态采样

如果你的数据集可以全量放入内存(比如用numpy数组存储),这种方式更灵活,每次调用输入函数都会重新采样负样本:

  • 核心思路:先筛选出正、负样本的索引,然后随机采样等量的负样本,合并后返回给Estimator
  • 代码示例:
import numpy as np
import tensorflow as tf

# 假设X是特征矩阵,y是标签数组(提前加载到内存)
X = np.load('features.npy')
y = np.load('labels.npy')

def balanced_input_fn():
    # 筛选正、负样本索引
    pos_indices = np.where(y == 1)[0]
    neg_indices = np.where(y == 0)[0]

    # 随机采样负样本,数量和正样本一致
    sampled_neg_indices = np.random.choice(neg_indices, size=len(pos_indices), replace=False)

    # 合并并打乱索引
    balanced_indices = np.concatenate([pos_indices, sampled_neg_indices])
    np.random.shuffle(balanced_indices)

    # 返回特征和标签,符合Estimator的输入格式
    return {'feature_column': X[balanced_indices]}, y[balanced_indices]
  • 优点:实现简单,无需修改数据流管道,每次训练都能拿到不同的负样本组合,降低过拟合风险
  • 注意点:如果数据量极大,全量加载到内存会导致OOM,这种情况下优先选方案一;另外要确保采样是无放回的,避免重复使用同一个负样本。

额外小建议(针对入门者)

  • 别只盯着准确率:不均衡数据集下,模型哪怕全预测负样本也能拿到99.9%的准确率,所以要重点监控召回率(Recall)和F1值,这两个指标能真实反映模型对正样本的识别能力
  • 可以试试加权损失:如果不想修改数据集,也可以给正样本设置更高的损失权重。预定义Estimator可以通过weight_column参数实现——把样本权重作为一个特征传入,正样本权重设为999(对应0.1%的占比),负样本设为1,这样模型会更重视正样本的预测
  • 对比实验:训练时可以同时跑“均衡数据集”和“原始数据集”的对比,直观感受下采样对模型性能的提升

如果你的两种预处理思路和上面的有出入,或者在代码实现上碰到具体问题,随时把细节贴出来,我再帮你调整~

内容的提问来源于stack exchange,提问作者Gianluca Micchi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 09:21:07