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

如何使用tf.data从50000条样本数据集中均匀采样2000个元素

实现方案

要实现从50000条样本中均匀随机采样2000个元素,可以通过tf.data.Dataset的内置打乱+截取接口实现,修改后的代码如下:

import tensorflow as tf

# 总样本量与采样数量
TOTAL_SAMPLES = 50000
SAMPLE_SIZE = 2000
bs = 32 # 替换为你实际使用的批次大小

dataset = tf.data.TFRecordDataset(path_filename_records)
dataset = (dataset
           .shuffle(buffer_size=TOTAL_SAMPLES, seed=42) # 全局打乱保证均匀采样,seed可固定保证结果可复现
           .take(SAMPLE_SIZE) # 截取前2000条打乱后的样本,即为均匀随机采样结果
           .map(parse_record, num_parallel_calls=tf.data.experimental.AUTOTUNE)
           .batch(bs)
           .prefetch(tf.data.experimental.AUTOTUNE)
          )

注意事项

  • shuffle的buffer_size必须设置为大于等于总样本数,才能保证全局完全打乱,采样结果符合均匀随机要求;如果内存不足以加载全部样本做buffer,可以设置buffer_size为总样本数的1/10以上,均匀性也能满足绝大多数场景需求
  • 把take操作放在map之前可以减少不必要的计算:仅对采样到的2000条样本做解析处理,运行性能更高

如果不想提前固定总样本数,也可以用随机过滤的方案实现:

import tensorflow as tf

SAMPLE_RATE = 2000 / 50000 # 采样率为0.04

dataset = tf.data.TFRecordDataset(path_filename_records)
dataset = (dataset
           .filter(lambda x: tf.random.uniform(()) < SAMPLE_RATE)
           .map(parse_record, num_parallel_calls=tf.data.experimental.AUTOTUNE)
           .batch(bs)
           .prefetch(tf.data.experimental.AUTOTUNE)
          )

该方案的采样数量会存在小范围波动,不需要提前知道总样本量,适合总样本数不确定的场景。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 18:39:03