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个流派类别)。
错误原因
- 目录遍历逻辑问题:
save_mfcc函数中使用os.walk遍历数据集目录时,会把所有子目录(包括数据集里的非流派目录,比如可能存在的images、features等辅助目录)都当作流派目录处理,导致标签值i-1不断递增,超过了10个流派对应的0-9范围。 - 标签生成方式不合理:依赖遍历的索引
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
相关产品推荐
相关产品推荐

