训练时如何从5个TFRecords文件中均衡采样数据?
当然有办法解决这个问题!在TensorFlow里,我们可以通过两种思路实现这种均衡采样的需求,一种是随机长期均衡,另一种是严格单批次均衡,你可以根据自己的训练需求来选:
方法一:随机长期均衡采样(灵活高效)
这种方法会让TensorFlow从每个TFRecords文件中按均等权重随机采样,长期来看每个文件的样本占比完全一致,单个batch可能会有微小波动,但整体符合均衡要求,实现起来也最简单:
- 首先定义你的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']
- 为每个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)
- 使用
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个样本,没有任何波动,可以用这种方法:
- 同样先完成解析函数和子数据集的创建(和方法一的前两步一致);
- 给每个子数据集先设置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
相关产品推荐
相关产品推荐

