如何高效筛选指定数量样本并合并为单个tf.data.Dataset?
问题描述
我有一个包含400多万条记录的大型TFRecord文件,数据集存在严重不平衡问题,部分标签的样本数量远多于其他标签。我希望筛选部分标签的指定数量样本以构建平衡数据集,但当前实现方法从33个标签各筛选1000条样本耗时超过24小时,以下是我的尝试代码:
import tensorflow as tf tf.compat.as_str( bytes_or_text='str', encoding='utf-8' ) try: tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect() print("Device:", tpu.master()) strategy = tf.distribute.TPUStrategy(tpu) except: strategy = tf.distribute.get_strategy() print("Number of replicas:", strategy.num_replicas_in_sync) ignore_order = tf.data.Options() ignore_order.experimental_deterministic = False dataset = tf.data.TFRecordDataset('/test.tfrecord') dataset = dataset.with_options(ignore_order) features, feature_lists = detect_schema(dataset) # 解码TFRecord序列化数据 def decode_data(serialized): X, y = tf.io.parse_single_sequence_example( serialized, context_features=features, sequence_features=feature_lists) return X['title'], y['subject'] dataset = dataset.map(lambda x: tf.py_function(func=decode_data, inp=[x], Tout=(tf.string, tf.string))) # 筛选并合并样本 def balanced_dataset(dataset, labels_list, sample_size=1000): datasets_list = [] for label in labels_list: # 筛选指定标签 locals()[label] = dataset.filter(lambda x, y: tf.greater(tf.reduce_sum(tf.cast(tf.equal(tf.constant(label, dtype=tf.int64), y), tf.float32)), tf.constant(0.))) # 添加指定数量的样本 datasets_list.append(locals()[label].take(sample_size)) concat_dataset = datasets_list[0] # 合并数据集 for dset in datasets_list[1:]: concat_dataset = concat_dataset.concatenate(dset) return concat_dataset balanced_data = balanced_dataset(tabledataset, labels_list=list(decod_dic.values()), sample_size=1000)
优化方案
核心问题分析
- 重复遍历数据集:当前代码对每个标签都单独遍历整个TFRecord文件做筛选,33个标签就要遍历33次400万条数据,这是性能瓶颈的核心来源。
tf.py_function的额外开销:该API会绕过TensorFlow的图优化机制,导致解码过程无法被高效编译执行。- 冗余的标签判断逻辑:用
tf.reduce_sum+类型转换的方式判断标签匹配,完全可以简化。
具体优化步骤
- 单次遍历完成所有标签采样:遍历数据集一次,同时为每个目标标签收集指定数量的样本,收集完成后立即停止对该标签的采样。
- 替换
tf.py_function为纯TensorFlow操作:直接在map中使用tf.io.parse_single_sequence_example,避免Python函数调用的额外开销。 - 简化标签匹配逻辑:用
tf.reduce_any替代求和判断,逻辑更简洁高效。
优化后代码
import tensorflow as tf import numpy as np try: tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect() print("Device:", tpu.master()) strategy = tf.distribute.TPUStrategy(tpu) except: strategy = tf.distribute.get_strategy() print("Number of replicas:", strategy.num_replicas_in_sync) ignore_order = tf.data.Options() ignore_order.experimental_deterministic = False dataset = tf.data.TFRecordDataset('/test.tfrecord') dataset = dataset.with_options(ignore_order) features, feature_lists = detect_schema(dataset) # 纯TensorFlow解码,移除tf.py_function def decode_data(serialized): X, y = tf.io.parse_single_sequence_example( serialized, context_features=features, sequence_features=feature_lists) return X['title'], y['subject'] # 启用并行解码,提升效率 dataset = dataset.map(decode_data, num_parallel_calls=tf.data.AUTOTUNE) def balanced_dataset(dataset, target_labels, sample_size=1000): # 初始化标签计数器与样本缓存 label_counters = {label: 0 for label in target_labels} collected_samples = {label: [] for label in target_labels} # 单次遍历数据集,同时收集所有目标标签的样本 for title, subject in dataset.as_numpy_iterator(): # 根据实际数据类型转换标签格式 label = subject.item() if isinstance(subject, np.ndarray) else subject if label in label_counters and label_counters[label] < sample_size: collected_samples[label].append((title, subject)) label_counters[label] += 1 # 所有标签收集完成后提前终止遍历 if all(count >= sample_size for count in label_counters.values()): break # 合并所有收集到的样本为tf.data.Dataset balanced_ds = None for label in target_labels: samples = collected_samples[label] ds = tf.data.Dataset.from_tensor_slices( ([s[0] for s in samples], [s[1] for s in samples]) ) balanced_ds = ds if balanced_ds is None else balanced_ds.concatenate(ds) return balanced_ds.shuffle(buffer_size=len(target_labels)*sample_size) # 修正原代码中的变量名错误(tabledataset改为dataset) balanced_data = balanced_dataset(dataset, labels_list=list(decod_dic.values()), sample_size=1000)
额外性能提升建议
- 添加预取操作:在
map后追加dataset = dataset.prefetch(tf.data.AUTOTUNE),实现数据加载与后续处理并行执行。 - 批量解码优化:如果业务允许,使用
tf.io.parse_sequence_example处理批量数据,进一步提升解码效率。 - 预建标签索引:若需要频繁按标签筛选数据,可以预先为TFRecord文件建立标签-数据位置的索引,后续直接通过索引定位目标数据。
内容的提问来源于stack exchange,提问作者Marlon Teixeira
相关产品推荐
相关产品推荐

