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

如何用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 21:40:23