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

Kotlin/Java环境下为yamnet.tflite准备.wav/.amr音频输入的方案求助

YAMNet TFLite 音频预处理正确实现方案

一、现有代码问题说明

现有代码存在三个核心问题会导致模型输出错误:

  • 硬编码44字节作为WAV头偏移量不通用:部分带自定义元数据的WAV头长度大于44字节,直接截取会把部分头数据算入PCM,导致输入错误
  • 未校验WAV参数:YAMNet要求输入必须是16kHz采样率、单声道、16bit位深的PCM数据,不符合要求的输入会直接导致分类错误
  • 切片逻辑缺陷:长度不足15600的末尾音频片段会被直接丢弃,短音频场景下甚至会返回空输入列表

二、WAV格式正确预处理方案

处理流程

  1. 解析WAV头,校验格式参数,动态定位PCM数据起始位置
  2. 将16bit有符号PCM数据转换为[-1.0, 1.0]区间的浮点数组
  3. 音频切片时不足15600采样点的末尾片段自动补0,不丢弃有效数据

修正后代码实现

import java.io.BufferedInputStream
import java.io.DataInputStream
import java.io.File
import java.io.FileInputStream
import java.nio.ByteBuffer
import java.nio.ByteOrder

object AudioConverter {
    // YAMNet 固定入参要求
    private const val TARGET_SAMPLE_RATE = 16000
    private const val INPUT_SAMPLE_LENGTH = 15600 // 对应0.975s音频@16kHz
    private const val MAX_16BIT_PCM = 32768.0f

    data class WavInfo(
        val channel: Int,
        val sampleRate: Int,
        val bitDepth: Int,
        val pcmStartOffset: Int,
        val pcmData: ByteArray
    )

    fun readWavForYamnet(path: File, step: Int = 0): List<FloatArray> {
        val wavInfo = parseWavFile(path)
        // 参数校验,不符合要求的需先做重采样/声道转换
        require(wavInfo.sampleRate == TARGET_SAMPLE_RATE) { "仅支持16kHz采样率的WAV文件" }
        require(wavInfo.channel == 1) { "仅支持单声道WAV文件" }
        require(wavInfo.bitDepth == 16) { "仅支持16bit位深的WAV文件" }

        // PCM转[-1.0,1.0]浮点数组
        val shortPcm = ByteBuffer.wrap(wavInfo.pcmData).order(ByteOrder.LITTLE_ENDIAN).asShortBuffer().let {
            val arr = ShortArray(it.remaining())
            it.get(arr)
            arr
        }
        val floatAudio = shortPcm.map { it.toFloat() / MAX_16BIT_PCM }.toFloatArray()

        // 切片+不足补0
        return sliceAudioWithPadding(floatAudio, INPUT_SAMPLE_LENGTH, step)
    }

    private fun parseWavFile(path: File): WavInfo {
        val fileBytes = path.readBytes()
        var offset = 0
        var channel = 1
        var sampleRate = TARGET_SAMPLE_RATE
        var bitDepth = 16

        // 校验RIFF头
        val riffId = String(fileBytes.sliceArray(offset until offset+4))
        require(riffId == "RIFF") { "不是合法WAV文件" }
        offset += 8
        val waveId = String(fileBytes.sliceArray(offset until offset+4))
        require(waveId == "WAVE") { "不是合法WAV文件" }
        offset +=4

        // 遍历块定位fmt和data段
        while (offset < fileBytes.size) {
            val blockId = String(fileBytes.sliceArray(offset until offset+4))
            val blockSize = ByteBuffer.wrap(fileBytes, offset+4, 4).order(ByteOrder.LITTLE_ENDIAN).int
            if (blockId == "fmt ") {
                val audioFormat = ByteBuffer.wrap(fileBytes, offset+8, 2).order(ByteOrder.LITTLE_ENDIAN).short
                require(audioFormat == 1.toShort()) { "仅支持PCM编码WAV文件" }
                channel = ByteBuffer.wrap(fileBytes, offset+10, 2).order(ByteOrder.LITTLE_ENDIAN).short.toInt()
                sampleRate = ByteBuffer.wrap(fileBytes, offset+12,4).order(ByteOrder.LITTLE_ENDIAN).int
                bitDepth = ByteBuffer.wrap(fileBytes, offset+22,2).order(ByteOrder.LITTLE_ENDIAN).short.toInt()
                offset += 8 + blockSize
                continue
            }
            if (blockId == "data") {
                val pcmData = fileBytes.sliceArray(offset+8 until offset+8+blockSize)
                return WavInfo(
                    channel = channel,
                    sampleRate = sampleRate,
                    bitDepth = bitDepth,
                    pcmStartOffset = offset+8,
                    pcmData = pcmData
                )
            }
            offset += 8 + blockSize
        }
        throw IllegalArgumentException("未找到WAV文件的PCM数据块")
    }

    private fun sliceAudioWithPadding(audio: FloatArray, sliceLength: Int, step: Int = 0): List<FloatArray> {
        val slices = mutableListOf<FloatArray>()
        val stepSize = if (step > 0) (sliceLength * (1f / (2 * step))).toInt() else sliceLength
        var start = 0
        while (start < audio.size) {
            val end = start + sliceLength
            val slice = if (end <= audio.size) {
                audio.copyOfRange(start, end)
            } else {
                audio.copyOfRange(start, audio.size) + FloatArray(end - audio.size)
            }
            slices.add(slice)
            start += stepSize
        }
        return slices
    }
}

调用方式

// 无滑窗切片
AudioConverter.readWavForYamnet(file).forEach { 
    audioTensor.load(it)
    classifier.classify(audioTensor)
}

// 带滑窗 step=2代表重叠率50%
AudioConverter.readWavForYamnet(file, step = 2).forEach { 
    audioTensor.load(it)
    classifier.classify(audioTensor)
}

三、AMR格式预处理方案

AMR是压缩音频格式,需要先解码再处理:

  • 解码:使用Android系统MediaCodec或者FFmpeg库,将AMR格式解码为PCM原始数据
  • 格式转换:AMR默认采样率为8kHz,解码后需要重采样到16kHz,合并为单声道,转换为16bit位深
  • 后续处理:转换后的PCM数据直接走上述WAV预处理的转浮点、切片流程即可输入模型

内容的提问来源于stack exchange,提问作者Rufan Khokhar

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 13:15:02