使用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
相关产品推荐
相关产品推荐

