Android集成TensorFlow Lite猫狗分类模型遇初始化错误求助
问题背景
我在Android应用中集成TensorFlow Lite模型实现猫狗分类,流程为:从Kaggle下载数据集,通过Teachable Machine训练并导出量化后的TFLite模型,集成代码运行时出现初始化错误。
错误日志
Error getting native address of native library: task_vision_jni_gms
java.lang.IllegalArgumentException: Error occurred when initializing ObjectDetector: Mobile SSD models are expected to have exactly 4 outputs, found 1
at org.tensorflow.lite.task.gms.vision.detector.ObjectDetector.initJniWithModelFdAndOptions(Native Method)
at org.tensorflow.lite.task.gms.vision.detector.ObjectDetector.zzb(Unknown Source:0)
at org.tensorflow.lite.task.gms.vision.detector.zzb.createHandle(org.tensorflow:tensorflow-lite-task-vision-play-services@@0.4.2:4)
at org.tensorflow.lite.task.core.TaskJniUtils$1.createHandle(TaskJniUtils.java:70)
at org.tensorflow.lite.task.core.TaskJniUtils.createHandleFromLibrary(TaskJniUtils.java:91)
at org.tensorflow.lite.task.core.TaskJniUtils.createHandleFromFdAndOptions(TaskJniUtils.java:66)
at org.tensorflow.lite.task.gms.vision.detector.ObjectDetector.createFromFileAndOptions(org.tensorflow:tensorflow-lite-task-vision-play-services@@0.4.2:2)
at com.affinidi.tfdemoone.ObjectDetectorHelper.setupObjectDetector(ObjectDetectorHelper.kt:104)
at com.affinidi.tfdemoone.ObjectDetectorHelper.detect(ObjectDetectorHelper.kt:121)
at com.affinidi.tfdemoone.MainActivity.detectObjects(MainActivity.kt:89)
at com.affinidi.tfdemoone.MainActivity.startCamera$lambda$4$lambda$3$lambda$2(MainActivity.kt:132)
at com.affinidi.tfdemoone.MainActivity.$r8$lambda$cwS3iJ069sufgGf-nT7H81EEGtQ(Unknown Source:0)
at com.affinidi.tfdemoone.MainActivity$$ExternalSyntheticLambda3.analyze(Unknown Source:2)
at androidx.camera.core.ImageAnalysis.lambda$setAnalyzer$2(ImageAnalysis.java:481)
at androidx.camera.core.ImageAnalysis$$ExternalSyntheticLambda2.analyze(Unknown Source:2)
at androidx.camera.core.ImageAnalysisAbstractAnalyzer.lambda$analyzeImage$0$androidx-camera-core-ImageAnalysisAbstractAnalyzer(ImageAnalysisAbstractAnalyzer.java:286)
at androidx.camera.core.ImageAnalysisAbstractAnalyzer$$ExternalSyntheticLambda1.run(Unknown Source:14)
at java.util.concurrent.ThreadPoolExecutor.runWorker(ThreadPoolExecutor.java:1167)
at java.util.concurrent.ThreadPoolExecutor$Worker.run(ThreadPoolExecutor.java:641)
at java.lang.Thread.run(Thread.java:920)
解决方案
核心原因
Teachable Machine导出的是图像分类模型(仅1个输出,对应分类概率),但代码中使用的ObjectDetector是专门为目标检测模型(如Mobile SSD,需要4个输出:边界框坐标、类别、置信度等)设计的API,两者类型不匹配导致错误。
修复步骤
替换API为图像分类专用的
ImageClassifier
修改ObjectDetectorHelper类,适配图像分类场景:class ObjectDetectorHelper( var threshold: Float = 0.5f, var numThreads: Int = 2, var maxResults: Int = 1, var currentDelegate: Int = 0, var currentModel: Int = 0, val context: Context, val objectDetectorListener: DetectorListener ) { private val TAG = "ImageClassificationHelper" private var imageClassifier: ImageClassifier? = null init { TfLiteGpu.isGpuDelegateAvailable(context).onSuccessTask { gpuAvailable: Boolean -> val optionsBuilder = TfLiteInitializationOptions.builder() if (gpuAvailable) { optionsBuilder.setEnableGpuDelegateSupport(true) } TfLiteVision.initialize(context, optionsBuilder.build()) }.addOnSuccessListener { objectDetectorListener.onInitialized() }.addOnFailureListener{ objectDetectorListener.onError("TfLiteVision failed to initialize: " + it.message) } } fun clearClassifier() { imageClassifier = null } private fun setupImageClassifier() { if (!TfLiteVision.isInitialized()) { Log.e(TAG, "setupImageClassifier: TfLiteVision is not initialized yet") return } val optionsBuilder = ImageClassifier.ImageClassifierOptions.builder() .setScoreThreshold(threshold) .setMaxResults(maxResults) val baseOptionsBuilder = BaseOptions.builder().setNumThreads(numThreads) when (currentDelegate) { DELEGATE_CPU -> {} DELEGATE_GPU -> baseOptionsBuilder.useGpu() DELEGATE_NNAPI -> baseOptionsBuilder.useNnapi() } optionsBuilder.setBaseOptions(baseOptionsBuilder.build()) val modelName = "model.tflite" try { imageClassifier = ImageClassifier.createFromFileAndOptions(context, modelName, optionsBuilder.build()) } catch (e: Exception) { objectDetectorListener.onError("Image classifier failed to initialize. See error logs for details") Log.e(TAG, "TFLite failed to load model with error: " + e.message) } } fun classify(image: Bitmap, imageRotation: Int) { if (!TfLiteVision.isInitialized()) { Log.e(TAG, "classify: TfLiteVision is not initialized yet") return } if (imageClassifier == null) { setupImageClassifier() } var inferenceTime = SystemClock.uptimeMillis() val imageProcessor = ImageProcessor.Builder().add(Rot90Op(-imageRotation / 90)).build() val tensorImage = imageProcessor.process(TensorImage.fromBitmap(image)) val results = imageClassifier?.classify(tensorImage) inferenceTime = SystemClock.uptimeMillis() - inferenceTime objectDetectorListener.onResults(results, inferenceTime, tensorImage.height, tensorImage.width) } interface DetectorListener { fun onInitialized() fun onError(error: String) fun onResults( results: MutableList<Classifications>?, inferenceTime: Long, imageHeight: Int, imageWidth: Int ) } companion object { const val DELEGATE_CPU = 0 const val DELEGATE_GPU = 1 const val DELEGATE_NNAPI = 2 } }修改MainActivity中的回调与调用逻辑
// 适配Classifications类型的结果回调 override fun onResults( results: MutableList<Classifications>?, inferenceTime: Long, imageHeight: Int, imageWidth: Int ) { runOnUiThread { val topCategory = results?.firstOrNull()?.categories?.firstOrNull() Log.i("resultssss", "分类结果: ${topCategory?.label}, 置信度: ${topCategory?.score}") } } // 修改detectObjects中的方法调用 private fun detectObjects(image: ImageProxy) { Log.i("resultssss", "5") image.use { bitmapBuffer.copyPixelsFromBuffer(image.planes[0].buffer) } Log.i("resultssss", "6") val imageRotation = image.imageInfo.rotationDegrees Log.i("resultssss", "7") objectDetectorHelper.classify(bitmapBuffer, imageRotation) Log.i("resultssss", "8") }确认依赖配置
确保build.gradle中包含图像分类所需依赖:implementation 'org.tensorflow:tensorflow-lite-task-vision-play-services:0.4.2'
额外说明
如果需求是目标检测(需要识别猫狗的位置),则需重新通过Teachable Machine训练目标检测模型,或更换为Mobile SSD类的目标检测模型,再使用ObjectDetector加载。
内容的提问来源于stack exchange,提问作者BraveEvidence

