Android端TensorFlow Lite模型输入报错:张量字节数不匹配
解决Android端TFLite模型输入张量不匹配问题
问题描述
在Android Studio 7.2.1中尝试给64×64图像的.tflite模型输入数据时触发错误:
E/AndroidRuntime: FATAL EXCEPTION: main Process: com.example.tag22, PID: 15184 java.lang.IllegalArgumentException: Cannot copy to a TensorFlowLite tensor (input_values:0) with 0 bytes from a Java Buffer with 49152 bytes.
Android端尝试的输入代码:
val catBitmap = getBitmapFromAsset("bwr.jpg")// a 64*64 image val output2 = TensorImage(DataType.FLOAT32) for (y in 0..63) { for (x in 0..63) { val px: Int = catBitmap.getPixel(x, y) // Get channel values from the pixel value. val r: Int = Color.red(px) val g: Int = Color.green(px) val b: Int = Color.blue(px) // Normalize channel values to [-1.0, 1.0]. This requirement depends // on the model. For example, some models might require values to be // normalized to the range [0.0, 1.0] instead. val rf = (r - 127) / 255.0f val gf = (g - 127) / 255.0f val bf = (b - 127) / 255.0f input.putFloat(rf) input.putFloat(gf) input.putFloat(bf) } } tflite.run(input, output2)
而Python端使用以下代码可成功运行:
image_filename='img.jpeg' input_data = tf.compat.v1.gfile.FastGFile(image_filename, 'rb').read() cc=[input_data ] input_data = np.array([input_data ]) interpreter.set_tensor(input_details[0]['index'], input_data)
问题分析
从Python代码可以明确,该模型的输入并非常规的归一化RGB浮点数组,而是直接接收图片文件的原始二进制字节数据。Android端代码错误地将图片解析为RGB通道并转换为float值,导致输入数据的格式、类型与模型要求完全不匹配,触发张量复制错误。
解决方案
在Android端直接读取图片的原始字节数据,按照模型要求的输入形状和类型填充到输入张量中,具体步骤如下:
修正后的Android代码示例
// 1. 读取assets中的图片原始字节数据 val inputStream = assets.open("bwr.jpg") val byteArray = inputStream.readBytes() inputStream.close() // 2. 获取模型输入张量并检查类型 val inputTensor = tflite.getInputTensor(0) val inputDataType = inputTensor.dataType() // 3. 根据张量类型填充对应格式的数据 when (inputDataType) { DataType.UINT8 -> { // 若输入为UINT8类型,直接写入字节数组 inputTensor.loadBuffer(byteArray) } DataType.STRING -> { // 若输入为STRING类型,将字节数组包装为ByteBuffer写入 val buffer = ByteBuffer.allocateDirect(byteArray.size) buffer.put(byteArray) buffer.rewind() inputTensor.loadBuffer(buffer) } else -> { throw IllegalArgumentException("不支持的输入张量类型: $inputDataType") } } // 4. 准备输出并运行模型 val outputTensor = tflite.getOutputTensor(0) val outputBuffer = ByteBuffer.allocateDirect(outputTensor.numBytes()) tflite.run(arrayOf(inputTensor.buffer), arrayOf(outputBuffer)) // 解析输出数据(示例:假设输出为float数组) outputBuffer.rewind() val outputArray = FloatArray(outputTensor.numElements()) outputBuffer.asFloatBuffer().get(outputArray)
额外注意事项
- 确保读取的图片文件与Python端使用的完全一致(尺寸、格式无差异)
- 可通过
inputTensor.shape()查看模型输入的形状,确保输入数据维度匹配 - 若模型要求批量输入(如Python中的
np.array([input_data])),需将字节数组包装为批量格式(例如创建包含单个字节数组的数组)
内容的提问来源于stack exchange,提问作者creative_sh
相关产品推荐
相关产品推荐

