TFLite语义分割模型Python端运行正常,Android端推理异常求助
问题排查与修复方案
以下是你代码中存在的可导致推理异常的问题及对应修复方法:
核心问题列表
- 输入缓冲区分配错误
TFLite 要求输入输出缓冲区必须使用直接内存(堆外内存),你当前代码中输入缓冲区使用ByteBuffer.allocate()分配堆内存,会导致推理输入读取异常,需改为allocateDirect()。 - 输入维度顺序错误
你注释标注输入形状为1 x 80(width) x 160(height) x 3,但绝大多数语义分割模型输入为批次 x 高度 x 宽度 x 通道(NHWC)顺序,结合你Colab正常运行的逻辑,实际输入应为1 x 160(height) x 80(width) x 3,你当前代码中循环遍历宽高的顺序、Bitmap缩放参数的宽高顺序都与模型要求不匹配,导致输入数据错乱。 - 输出转Bitmap逻辑错误
你当前仅将输出值赋值给红色通道,灰度图需要三个通道数值一致,否则会出现颜色异常;同时如果输出维度顺序为高度x宽度,你按一维数组直接读取的顺序也会对应错误,导致输出图像错乱。 - 重复初始化Interpreter
每次调用detect方法都重新加载模型、新建Interpreter实例,会造成严重的性能损耗、内存泄漏,甚至可能出现推理异常,建议将Interpreter初始化逻辑移到方法外,全局复用一个实例。
修复后的核心代码
// 建议将interpreter改为全局变量,初始化一次即可 private lateinit var interpreter: Interpreter // 在Activity/Fragment初始化时调用一次即可 fun initInterpreter(assets: AssetManager) { val modelFile = loadModelFile(assets, MODEL_FILENAME) interpreter = Interpreter(modelFile) } fun detect(bitmap : Bitmap?) : Bitmap? { if (bitmap == null) { return null } val h = bitmap.height val w = bitmap.width val outputBitmap : Bitmap = Bitmap.createBitmap(80, 160, Bitmap.Config.ARGB_8888) // 输出形状修正为NHWC顺序:1 x 160(height) x 80(width) x 1 val outputByteBuffer: ByteBuffer = ByteBuffer.allocateDirect(1 * 160 * 80 * 1 * 4) outputByteBuffer.order(ByteOrder.nativeOrder()) // 输入形状修正为NHWC顺序:1 x 160(height) x 80(width) x 3 val inputByteBuffer: ByteBuffer = ByteBuffer.allocateDirect(1* 160 * 80 * 3 * 4) inputByteBuffer.order(ByteOrder.nativeOrder()) // 缩放bitmap宽高对应模型输入的宽度、高度,修正为160高、80宽 val inputBitmap = Bitmap.createScaledBitmap(bitmap,80,160,false) inputByteBuffer.rewind() outputByteBuffer.rewind() val IMAGE_MEAN: Float = 128f val IMAGE_STD: Float = 128f val inputPixelValues = IntArray(80 * 160) inputBitmap.getPixels(inputPixelValues, 0, 80, 0, 0, 80, 160) var pixel = 0 // 循环顺序修正为先行(高度)后列(宽度),匹配NHWC顺序 for (i in 0 until 160) { for (j in 0 until 80) { val `val` = inputPixelValues[pixel++] inputByteBuffer.putFloat(((`val` shr 16 and 0xFF) - IMAGE_MEAN) / IMAGE_STD) inputByteBuffer.putFloat(((`val` shr 8 and 0xFF) - IMAGE_MEAN) / IMAGE_STD) inputByteBuffer.putFloat(((`val` and 0xFF) - IMAGE_MEAN) / IMAGE_STD) } } inputBitmap.recycle() // 运行推理 interpreter.run(inputByteBuffer, outputByteBuffer) // 转换输出为Bitmap outputByteBuffer.rewind() val pixels = IntArray(80 * 160) for (i in 0 until 80 * 160) { val a = 0xFF val gray: Int = (outputByteBuffer.float * 255.0f).toInt().coerceIn(0, 255) // 灰度图三个通道赋值相同 pixels[i] = a shl 24 or (gray shl 16) or (gray shl 8) or gray } outputBitmap.setPixels(pixels, 0, 80, 0, 0, 80, 160) // 缩放回原尺寸 val resizedOutputBitmap = Bitmap.createScaledBitmap(outputBitmap,w,h,false) outputBitmap.recycle() return resizedOutputBitmap }
内容的提问来源于stack exchange,提问作者amira yasmine
相关产品推荐
相关产品推荐

