Kotlin/Java环境下为yamnet.tflite准备.wav/.amr音频输入的方案求助
YAMNet TFLite 音频预处理正确实现方案
一、现有代码问题说明
现有代码存在三个核心问题会导致模型输出错误:
- 硬编码44字节作为WAV头偏移量不通用:部分带自定义元数据的WAV头长度大于44字节,直接截取会把部分头数据算入PCM,导致输入错误
- 未校验WAV参数:YAMNet要求输入必须是16kHz采样率、单声道、16bit位深的PCM数据,不符合要求的输入会直接导致分类错误
- 切片逻辑缺陷:长度不足15600的末尾音频片段会被直接丢弃,短音频场景下甚至会返回空输入列表
二、WAV格式正确预处理方案
处理流程
- 解析WAV头,校验格式参数,动态定位PCM数据起始位置
- 将16bit有符号PCM数据转换为[-1.0, 1.0]区间的浮点数组
- 音频切片时不足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
相关产品推荐
相关产品推荐

