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

训练时如何从5个TFRecords文件中均衡采样数据?

当然有办法解决这个问题!在TensorFlow里,我们可以通过两种思路实现这种均衡采样的需求,一种是随机长期均衡,另一种是严格单批次均衡,你可以根据自己的训练需求来选:

方法一:随机长期均衡采样(灵活高效)

这种方法会让TensorFlow从每个TFRecords文件中按均等权重随机采样,长期来看每个文件的样本占比完全一致,单个batch可能会有微小波动,但整体符合均衡要求,实现起来也最简单:

  1. 首先定义你的TFRecords解析函数,根据你实际存储的特征结构来编写:
def parse_tfrecord_fn(example):
    # 替换成你自己的特征描述
    feature_description = {
        'image': tf.io.FixedLenFeature([], tf.string),
        'label': tf.io.FixedLenFeature([], tf.int64),
        # 这里添加你所有的特征字段
    }
    example = tf.io.parse_single_example(example, feature_description)
    # 按需添加后续处理,比如图像解码、归一化等
    example['image'] = tf.io.decode_jpeg(example['image'], channels=3)
    example['image'] = tf.cast(example['image'], tf.float32) / 255.0
    return example['image'], example['label']
  1. 为每个TFRecords文件创建独立的数据集:
# 替换成你的5个TFRecords文件路径
tfrecord_paths = ['obj1.tfrecord', 'obj2.tfrecord', 'obj3.tfrecord', 'obj4.tfrecord', 'obj5.tfrecord']

# 逐个创建数据集并应用解析逻辑
datasets = []
for path in tfrecord_paths:
    ds = tf.data.TFRecordDataset(path)
    ds = ds.map(parse_tfrecord_fn, num_parallel_calls=tf.data.AUTOTUNE)
    ds = ds.repeat()  # 重复数据集,避免训练中途耗尽数据
    datasets.append(ds)
  1. 使用sample_from_datasets进行均衡采样,再设置batch size:
# 给每个数据集设置相同的权重,实现均衡采样
balanced_ds = tf.data.Dataset.sample_from_datasets(datasets, weights=[1.0]*5)
# 最后设置批量大小为50
balanced_ds = balanced_ds.batch(50).prefetch(tf.data.AUTOTUNE)

这样训练时,每个batch里来自每个文件的样本数量会大致是10个,长期来看完全均衡,适合大多数训练场景。

方法二:严格单批次均衡采样(精准分配)

如果你需要每个batch里严格从每个文件取10个样本,没有任何波动,可以用这种方法:

  1. 同样先完成解析函数和子数据集的创建(和方法一的前两步一致);
  2. 给每个子数据集先设置batch size为10,然后合并这些batch:
# 每个子数据集先batch成10个样本
batched_sub_ds = [ds.batch(10) for ds in datasets]
# 将这些batched数据集zip起来,每次会从每个子数据集取一个10样本的batch
zipped_ds = tf.data.Dataset.zip(tuple(batched_sub_ds))

# 定义合并函数,把5个10样本的batch合并成一个50样本的batch
def merge_batches(*batches):
    # 每个batch是(图像张量, 标签张量),将它们在维度0上拼接
    all_images = tf.concat([batch[0] for batch in batches], axis=0)
    all_labels = tf.concat([batch[1] for batch in batches], axis=0)
    return all_images, all_labels

# 应用合并函数,得到最终的均衡数据集
strict_balanced_ds = zipped_ds.map(merge_batches).prefetch(tf.data.AUTOTUNE)

这样每个输出的batch必然是50个样本,且每个原TFRecords文件恰好贡献10个,适合对数据分布有严格要求的场景。

一些额外注意点

  • 一定要给每个子数据集加上repeat(),否则当某个文件的样本被取完后,训练会直接中断;
  • 如果你的TFRecords文件样本数量差异很大,repeat()会自动循环补充该文件的样本,保证持续供应;
  • num_parallel_calls=tf.data.AUTOTUNE和prefetch(tf.data.AUTOTUNE)是为了加速数据加载,建议加上,提升训练效率。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 04:36:08