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.
相关产品推荐
相关产品推荐

