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

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.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 20:42:07