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

使用reduce计算TensorFlow数据集基数时触发exclude_short_tracks报错咨询

音频TF数据集计数问题解决

问题1:错误原因与exclude_short_tracks的调用环节

  • 错误根源:exclude_short_tracks定义需要path和label两个参数,但在filter的tf.py_function中只传入了[x](仅path),未传入y(label),导致调用时参数缺失。
  • 调用环节:TF Dataset的所有变换(filter、map等)都是惰性执行的,只有触发数据消费操作(如reduce、as_numpy_iterator()、迭代取数)时,才会从头到尾执行整个数据处理流水线。你用reduce计数时,会触发前面的filter步骤,进而调用exclude_short_tracks,此时参数不足就抛出了错误。

问题2:获取数据集基数的推荐方法

首先必须修复filter的参数问题,同时修正exclude_short_tracks的冗余参数和逻辑矛盾:

第一步:修复代码错误

def create_dataset(audio_paths, audio_classes):
    ds = tf.data.Dataset.zip(
        tf.data.Dataset.from_tensor_slices(audio_paths),
        tf.data.Dataset.from_tensor_slices(audio_classes)
    )
    # 修复:传入x和y两个参数给exclude_short_tracks
    ds = ds.filter(lambda x, y: tf.py_function(exclude_short_tracks, [x, y], tf.bool))
    ds = ds.map(lambda x, y: (tf.py_function(make_mel, [x], tf.float32), y))
    return ds

# 修复:去掉未使用的label参数,同时修正逻辑(原注释要求返回"文件长于SAMPLING_LENGTH"的结果)
def exclude_short_tracks(path, label):
    path = path.numpy().decode('ascii')  # 原代码的[0]需根据path张量实际结构调整
    length = librosa.get_duration(path = path)
    return length >= SAMPLING_LENGTH

第二步:推荐的计数方法

  • 方法1:reduce高效计数(无需加载全部数据到内存):
count = train_ds.reduce(0, lambda x, _: x + 1).numpy()
print("Training dataset cardinality:", count)
  • 方法2:迭代器统计:
count = len(list(train_ds.as_numpy_iterator()))
print("Training dataset cardinality:", count)
  • 方法3:提前预计算(避免重复触发数据集处理):
    在创建数据集前直接统计符合条件的样本数,后续可直接复用该数值:
def count_valid_samples(audio_paths):
    count = 0
    for path in audio_paths:
        length = librosa.get_duration(path=path)
        if length >= SAMPLING_LENGTH:
            count +=1
    return count

train_count = count_valid_samples(training_paths)
print("Training dataset cardinality:", train_count)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 10:14:54