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

如何高效筛选指定数量样本并合并为单个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)
优化方案

核心问题分析

  1. 重复遍历数据集:当前代码对每个标签都单独遍历整个TFRecord文件做筛选,33个标签就要遍历33次400万条数据,这是性能瓶颈的核心来源。
  2. tf.py_function的额外开销:该API会绕过TensorFlow的图优化机制,导致解码过程无法被高效编译执行。
  3. 冗余的标签判断逻辑:用tf.reduce_sum+类型转换的方式判断标签匹配,完全可以简化。

具体优化步骤

  1. 单次遍历完成所有标签采样:遍历数据集一次,同时为每个目标标签收集指定数量的样本,收集完成后立即停止对该标签的采样。
  2. 替换tf.py_function为纯TensorFlow操作:直接在map中使用tf.io.parse_single_sequence_example,避免Python函数调用的额外开销。
  3. 简化标签匹配逻辑:用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 08:55:23