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

Android端使用YAMNET TensorFlow Lite模型实现音频分类求助

Android 中使用 YAMNET TensorFlow Lite 实现声音分类完整流程

1. 项目依赖与模型准备

  • 在 app 模块的 build.gradle 中添加 TensorFlow Lite 相关依赖:
dependencies {
    implementation 'org.tensorflow:tensorflow-lite:2.15.0'
    implementation 'org.tensorflow:tensorflow-lite-support:0.4.4'
    implementation 'org.tensorflow:tensorflow-lite-metadata:0.4.4'
}
  • 下载 YAMNET 的 .tflite 模型,放置到 src/main/assets 目录;同时准备模型对应的类别标签文件(共521个音频类别),保存为 src/main/assets/yamnet_labels.txt。
  • 使用 Android Studio 的 TensorFlow Lite Model Binding 工具生成模型绑定类:右键点击 assets 中的 .tflite 文件,选择 TensorFlow Lite > Generate TensorFlow Lite Model Class,生成你代码中用到的 YamnetClassification 类。

2. 音频文件预处理

YAMNET 要求输入为 16kHz 采样率、单声道、16位 PCM 格式 的音频,需将数据归一化到 [-1.0, 1.0] 的 float32 数组,最终输入维度固定为 [15600](对应约0.975秒音频)。以下是预处理代码:

import android.media.MediaExtractor
import android.media.MediaFormat
import java.io.File
import java.nio.ByteBuffer
import java.nio.ByteOrder
import java.nio.ShortBuffer

fun preprocessWavFile(wavFile: File): ByteBuffer {
    val extractor = MediaExtractor()
    extractor.setDataSource(wavFile.path)

    // 定位音频轨道
    var audioTrackIndex = -1
    for (i in 0 until extractor.trackCount) {
        val format = extractor.getTrackFormat(i)
        if (format.getString(MediaFormat.KEY_MIME)?.startsWith("audio/") == true) {
            audioTrackIndex = i
            break
        }
    }
    require(audioTrackIndex != -1) { "未检测到音频轨道" }

    extractor.selectTrack(audioTrackIndex)
    val format = extractor.getTrackFormat(audioTrackIndex)
    val sampleRate = format.getInteger(MediaFormat.KEY_SAMPLE_RATE)
    require(sampleRate == 16000) { "YAMNET仅支持16kHz采样率音频" }

    val channelCount = format.getInteger(MediaFormat.KEY_CHANNEL_COUNT)
    require(channelCount == 1) { "YAMNET仅支持单声道音频" }

    // 读取音频数据为short数组(16位PCM)
    val shortBuffer = ShortBuffer.allocate(15600)
    val byteBuffer = ByteBuffer.allocateDirect(15600 * 2)
    byteBuffer.order(ByteOrder.nativeOrder())

    while (shortBuffer.hasRemaining() && extractor.readSampleData(byteBuffer, 0) != -1) {
        byteBuffer.flip()
        shortBuffer.put(byteBuffer.asShortBuffer())
        byteBuffer.clear()
        extractor.advance()
    }
    extractor.release()

    // 音频长度不足时补0
    while (shortBuffer.hasRemaining()) {
        shortBuffer.put(0.toShort())
    }
    shortBuffer.rewind()

    // 转换为float32并归一化到[-1.0, 1.0]
    val floatArray = FloatArray(15600)
    for (i in floatArray.indices) {
        floatArray[i] = shortBuffer.get(i) / 32768.0f
    }

    // 转换为模型所需的ByteBuffer
    val inputBuffer = ByteBuffer.allocateDirect(15600 * 4)
    inputBuffer.order(ByteOrder.nativeOrder())
    inputBuffer.asFloatBuffer().put(floatArray)
    inputBuffer.rewind()

    return inputBuffer
}

3. 模型推理与结果解析

结合你提供的代码,完善推理流程并解析分类结果:

import android.content.Context
import org.tensorflow.lite.support.tensorbuffer.TensorBuffer
import java.io.BufferedReader
import java.io.InputStreamReader

class YamnetAudioClassifier(private val context: Context) {
    private val model = YamnetClassification.newInstance(context)
    private val labels by lazy { loadLabels() }

    private fun loadLabels(): List<String> {
        return BufferedReader(InputStreamReader(context.assets.open("yamnet_labels.txt")))
            .readLines()
            .filter { it.isNotEmpty() }
    }

    fun classifyWavFile(wavFile: File): String {
        val inputBuffer = preprocessWavFile(wavFile)

        // 创建模型输入TensorBuffer
        val audioClip = TensorBuffer.createFixedSize(intArrayOf(15600), org.tensorflow.lite.support.common.DataType.FLOAT32)
        audioClip.loadBuffer(inputBuffer)

        // 执行推理
        val outputs = model.process(audioClip)
        val scores = outputs.scoresAsTensorBuffer.floatArray

        // 找到置信度最高的类别
        var maxScoreIndex = 0
        var maxScore = 0.0f
        scores.forEachIndexed { index, score ->
            if (score > maxScore) {
                maxScore = score
                maxScoreIndex = index
            }
        }

        return "类别:${labels[maxScoreIndex]},置信度:${String.format("%.2f", maxScore * 100)}%"
    }

    // 释放模型资源
    fun release() {
        model.close()
    }
}

4. 使用示例

在 Activity 或后台线程中调用分类器:

// 假设目标WAV文件已下载到本地存储
val wavFile = File(getExternalFilesDir(null), "cough.wav")
val classifier = YamnetAudioClassifier(this)
val result = classifier.classifyWavFile(wavFile)
Log.d("YamnetResult", result)
// 使用完毕后释放资源
classifier.release()

注意事项

  • 确保WAV文件为16kHz单声道格式,不符合需先做格式转换;
  • 读取外部存储音频时,需申请对应权限(Android 13+ 为 READ_MEDIA_AUDIO,低版本为 READ_EXTERNAL_STORAGE);
  • 模型推理需放在后台线程执行,避免阻塞主线程。

内容的提问来源于stack exchange,提问作者M. K.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 02:37:05