Android中基于YAMNet TFLite模型实现音频片段分类的方案问询
将YAMNet输入从实时录音改为本地音频片段的实现方案
YAMNet模型要求输入为16kHz采样率、单声道、16位PCM格式的音频,且需转换为归一化到[-1.0, 1.0]的FLOAT32张量(尺寸[15600])。下面是具体修改步骤:
一、完整实现代码
1. 读取本地音频并转换为模型要求的输入格式
val model = YamnetClassification.newInstance(this) try { // 1. 读取raw目录下的音频文件 val inputStream = resources.openRawResource(R.raw.s1) val audioBytes = inputStream.readBytes() inputStream.close() // 2. 将音频转换为16kHz单声道16位PCM格式 val pcmData = convertAudioTo16kHzMonoPCM(audioBytes) // 3. 将PCM数据转为模型所需的ByteBuffer val byteBuffer = ByteBuffer.allocateDirect(15600 * 4) // 每个float占4字节 byteBuffer.order(ByteOrder.nativeOrder()) val sampleCount = min(pcmData.size, 15600) for (i in 0 until sampleCount) { // 16位PCM转归一化float(范围-1.0到1.0) byteBuffer.putFloat(pcmData[i] / 32768.0f) } byteBuffer.rewind() // 4. 喂入模型执行推理 val audioClip = TensorBuffer.createFixedSize(intArrayOf(15600), DataType.FLOAT32) audioClip.loadBuffer(byteBuffer) val outputs = model.process(audioClip) val scores = outputs.scoresAsTensorBuffer // 这里可以处理scores结果,比如获取最高分类的索引和置信度 } finally { // 用完模型必须关闭,避免内存泄漏 model.close() }
2. 音频格式转换工具函数
这个函数负责将任意WAV音频转成16kHz单声道16位PCM格式:
@Throws(IOException::class) private fun convertAudioTo16kHzMonoPCM(audioBytes: ByteArray): ShortArray { val inputStream = ByteArrayInputStream(audioBytes) val extractor = MediaExtractor().apply { setDataSource(inputStream.fd) // 定位音频轨道 val audioTrackIdx = (0 until trackCount).firstOrNull { getTrackFormat(it).getString(MediaFormat.KEY_MIME)?.startsWith("audio/") == true } ?: throw IOException("无有效音频轨道") selectTrack(audioTrackIdx) } val inputFormat = extractor.getTrackFormat(extractor.selectedTrackIndex) val codec = MediaCodec.createDecoderByType(inputFormat.getString(MediaFormat.KEY_MIME)!!).apply { configure(inputFormat, null, null, 0) start() } val pcmList = mutableListOf<Short>() val bufferInfo = MediaCodec.BufferInfo() var isDone = false while (!isDone) { // 填充输入数据到解码器 val inputBufferId = codec.dequeueInputBuffer(10000) if (inputBufferId >= 0) { val inputBuffer = codec.getInputBuffer(inputBufferId)!! val sampleSize = extractor.readSampleData(inputBuffer, 0) if (sampleSize < 0) { codec.queueInputBuffer(inputBufferId, 0, 0, 0L, MediaCodec.BUFFER_FLAG_END_OF_STREAM) isDone = true } else { codec.queueInputBuffer(inputBufferId, 0, sampleSize, extractor.sampleTime, 0) extractor.advance() } } // 读取解码后的PCM数据 val outputBufferId = codec.dequeueOutputBuffer(bufferInfo, 10000) if (outputBufferId >= 0) { val outputBuffer = codec.getOutputBuffer(outputBufferId)!! outputBuffer.position(bufferInfo.offset) val pcmBuffer = ShortArray(bufferInfo.size / 2) // 16位PCM每个采样占2字节 outputBuffer.asShortBuffer().get(pcmBuffer) pcmList.addAll(pcmBuffer.toList()) codec.releaseOutputBuffer(outputBufferId, false) if (bufferInfo.flags and MediaCodec.BUFFER_FLAG_END_OF_STREAM != 0) { isDone = true } } } // 释放资源 codec.stop() codec.release() extractor.release() inputStream.close() // 重采样到16kHz(如果原音频采样率不是16kHz) return resamplePcm(pcmList.toShortArray(), inputFormat.getInteger(MediaFormat.KEY_SAMPLE_RATE), 16000) }
3. PCM重采样函数
如果原音频采样率不是16kHz,需要添加重采样逻辑:
private fun resamplePcm(inputPcm: ShortArray, originalSampleRate: Int, targetSampleRate: Int): ShortArray { if (originalSampleRate == targetSampleRate) return inputPcm val ratio = targetSampleRate.toFloat() / originalSampleRate val outputSize = (inputPcm.size * ratio).toInt() val outputPcm = ShortArray(outputSize) for (i in 0 until outputSize) { val inputIndex = (i / ratio).toInt() outputPcm[i] = inputPcm[min(inputIndex, inputPcm.size - 1)] } return outputPcm }
二、关键注意事项
- 输入格式要求:YAMNet只接受16kHz单声道16位PCM音频,任何其他格式都必须先转换,否则推理结果无效。
- 张量尺寸:模型输入必须是15600个采样点(对应约0.975秒音频),如果音频过长,截取前15600个采样点即可;如果过短,可以补零填充。
- 归一化处理:16位PCM的取值范围是-32768到32767,必须除以32768.0f转换为[-1.0, 1.0]的float值,这是模型的硬性要求。
内容的提问来源于stack exchange,提问作者M. K.
相关产品推荐
相关产品推荐

