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

TensorFlow训练音乐流派分类模型时InvalidArgumentError错误解决

音乐流派分类模型训练错误解决:标签超出有效范围

问题重现

使用TensorFlow训练基于GTZAN数据集的音乐流派分类模型时,运行model.fit()出现如下错误:

InvalidArgumentError: Graph execution error:
Node: 'sparse_categorical_crossentropy/SparseSoftmaxCrossEntropyWithLogits/SparseSoftmaxCrossEntropyWithLogits'
Received a label value of 12 which is outside the valid range of [0, 10).  Label values: 3 6 9 7 4 7 11 12

核心问题是:生成的标签值超过了模型输出层的类别范围(模型最后一层是Dense(10, activation="softmax"),对应0-9共10个流派类别)。

错误原因

  1. 目录遍历逻辑问题:save_mfcc函数中使用os.walk遍历数据集目录时,会把所有子目录(包括数据集里的非流派目录,比如可能存在的images、features等辅助目录)都当作流派目录处理,导致标签值i-1不断递增,超过了10个流派对应的0-9范围。
  2. 标签生成方式不合理:依赖遍历的索引i-1生成标签,没有考虑到可能存在的无效目录,导致标签出现10、11、12等超出范围的值。

解决方案

方案1:过滤无效目录,只保留GTZAN官方10个流派

GTZAN数据集的10个标准流派为:blues、classical、country、disco、hiphop、jazz、metal、pop、reggae、rock。修改save_mfcc函数,只处理这些目录:

def save_mfcc(dataset_path  , n_mfcc = 13 , n_fft = 2048 , hop_length = 512 , num_segments = 5):
    data = {
        'mapping' :[] ,
        'mfcc' :[] ,
        'labels' : []
    }
    # 定义有效流派列表
    VALID_GENRES = {'blues', 'classical', 'country', 'disco', 'hiphop', 'jazz', 'metal', 'pop', 'reggae', 'rock'}
    num_sample_per_segment = int(SAMPLES_PER_TRACK / num_segments)
    expected_num_mfcc_vector_per_segment = math.ceil(num_sample_per_segment / hop_length)
    print(f'{expected_num_mfcc_vector_per_segment} that is the length of the sequence')
    
    for i , (dirpath , dirnames , filenames) in enumerate(os.walk(dataset_path)):
        if dirpath != dataset_path:
            semantic_label = dirpath.split('/')[-1]
            # 跳过非有效流派的目录
            if semantic_label not in VALID_GENRES:
                print(f"跳过非流派目录: {semantic_label}")
                continue
            data['mapping'].append(semantic_label)
            for file in filenames:
                file_path = os.path.join(dirpath , file)
                try :
                    signal , sr = librosa.load(file_path , sr = SAMPLE_RATE)
                    for s in range(num_segments):
                        start_sample = num_sample_per_segment * s
                        finish_sample = start_sample + num_sample_per_segment
                        mfcc = librosa.feature.mfcc(y = signal[start_sample:finish_sample] , sr = SAMPLE_RATE , 
                                                   n_mfcc=13 , n_fft = n_fft , hop_length = hop_length)
                        mfcc = mfcc.T
                        if len(mfcc) == expected_num_mfcc_vector_per_segment:
                            data['mfcc'].append(mfcc.tolist())
                            # 用mapping的索引作为标签,保证0-9连续
                            data['labels'].append(len(data['mapping'])-1)
                except:
                    pass
            print(f"{semantic_label} is loaded successfully")
    return data

方案2:修正标签生成逻辑,确保标签连续有效

如果不想硬编码流派列表,可以修改标签生成方式,基于mapping列表的索引生成标签,确保标签始终是0到流派数量-1的连续值:

# 替换原代码中的data['labels'].append(i-1)
data['labels'].append(data['mapping'].index(semantic_label))

这样即使遍历到无效目录后跳过,标签依然会保持连续,不会出现断层或超出范围的值。

方案3:事后过滤无效标签样本

如果已经生成了数据,可以在划分数据集前过滤掉标签超出0-9范围的样本:

data = np.array(data_dict['mfcc'])
label = np.array(data_dict['labels']).reshape(-1 , 1)

# 过滤标签不在0-9范围内的样本
valid_mask = (label >= 0) & (label < 10)
# 提取有效样本
data = data[valid_mask.flatten()]
label = label[valid_mask.flatten()]

# 后续划分数据集的代码不变

验证

修改后,运行以下代码确认标签范围:

print("标签最小值:", np.min(label))
print("标签最大值:", np.max(label))

确认输出的最大值为9,最小值为0后,再执行模型训练即可解决错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 10:57:02