如何用Python运行UrbanSound8K训练的TFLite音频分类模型并处理外部音频
处理外部音频并运行TFLite音频分类模型
核心前提:预处理必须和训练一致
模型只接受和训练数据格式完全匹配的输入,处理外部音频时,必须复刻训练URBANSOUND8k时的预处理流程,关键点包括:
- 采样率:和训练时一致(URBANSOUND8k默认是22050Hz,以你实际训练参数为准)
- 音频时长:裁剪或补零到训练时使用的固定时长(比如1秒)
- 特征提取:使用相同的特征类型(如MFCC、梅尔频谱图)及参数(如MFCC系数个数、梅尔滤波器数量)
- 归一化:如果训练时对特征做了归一化,必须用训练集计算的均值/标准差做同样处理
完整代码实现
下面是Python环境下的可运行代码,涵盖音频预处理、TFLite模型加载与推理:
1. 依赖安装
先确保安装所需库:
pip install librosa tensorflow # 若仅需轻量运行,可替换tensorflow为 tflite_runtime
2. 预处理与推理代码
import librosa import numpy as np import tensorflow as tf # 若用tflite_runtime,替换上面的导入为: # from tflite_runtime.interpreter import Interpreter def preprocess_audio(audio_path, target_sr=22050, target_duration=1, n_mfcc=13): # 加载音频并转单通道 y, sr = librosa.load(audio_path, sr=target_sr, mono=True) # 统一音频长度到目标时长 target_samples = target_sr * target_duration if len(y) > target_samples: y = y[:target_samples] # 过长则裁剪 elif len(y) < target_samples: y = np.pad(y, (0, target_samples - len(y)), mode='constant') # 过短则补零 # 提取MFCC特征(和训练时一致,若用梅尔频谱图则替换为librosa.feature.melspectrogram) mfcc = librosa.feature.mfcc(y=y, sr=sr, n_mfcc=n_mfcc) # 调整形状匹配模型输入:根据你训练时的输入维度调整,示例为(1, n_mfcc, time_steps, 1) mfcc = np.expand_dims(mfcc, axis=0) # 增加batch维度 mfcc = np.expand_dims(mfcc, axis=-1) # 增加通道维度 # 若训练时做了特征归一化,添加以下步骤(替换为你训练集的均值和标准差) # mfcc_mean = 0.0 # 训练集MFCC均值 # mfcc_std = 1.0 # 训练集MFCC标准差 # mfcc = (mfcc - mfcc_mean) / mfcc_std return mfcc.astype(np.float32) def run_tflite_inference(model_path, processed_audio): # 加载TFLite模型 interpreter = tf.lite.Interpreter(model_path=model_path) interpreter.allocate_tensors() # 获取输入输出张量信息 input_details = interpreter.get_input_details() output_details = interpreter.get_output_details() # 检查输入形状是否匹配 assert processed_audio.shape == input_details[0]['shape'], \ f"输入形状不匹配:模型需要{input_details[0]['shape']},当前是{processed_audio.shape}" # 传入输入数据并推理 interpreter.set_tensor(input_details[0]['index'], processed_audio) interpreter.invoke() # 获取输出结果 output_probs = interpreter.get_tensor(output_details[0]['index'])[0] predicted_class_idx = np.argmax(output_probs) # URBANSOUND8k的类别标签 class_labels = [ "air_conditioner", "car_horn", "children_playing", "dog_bark", "drilling", "engine_idling", "gun_shot", "jackhammer", "siren", "street_music" ] return class_labels[predicted_class_idx], output_probs[predicted_class_idx] if __name__ == "__main__": # 替换为你的音频文件和TFLite模型路径 AUDIO_PATH = "test_audio.wav" MODEL_PATH = "urban_sound_model.tflite" # 预处理音频 processed_audio = preprocess_audio(AUDIO_PATH) # 运行推理 pred_label, pred_confidence = run_tflite_inference(MODEL_PATH, processed_audio) print(f"预测结果:{pred_label},置信度:{pred_confidence:.4f}")
关键注意事项
- 输入形状匹配:可以通过
input_details[0]['shape']查看模型要求的输入形状,调整预处理后的特征维度(比如有的模型输入是(1, time_steps, n_mfcc),需要转置MFCC的维度) - 特征类型匹配:如果训练时用的是梅尔频谱图、频谱图等其他特征,替换预处理中的特征提取函数即可
- 轻量运行:如果部署在资源受限设备上,使用
tflite_runtime代替完整的tensorflow库,减少内存占用
内容的提问来源于stack exchange,提问作者Onkar Mane
相关产品推荐
相关产品推荐

