如何将Google Colab中Keras训练的猫狗识别ML模型集成到Android Studio
核心解决方案:无需重写Java代码,直接导出Keras模型为TensorFlow Lite格式部署到Android
你完全可以直接导出Colab里训练好的模型使用,不需要重写模型的训练和推理逻辑,具体操作步骤如下:
步骤1:在Colab导出Keras模型为TFLite格式
Keras是TensorFlow官方的高阶API,直接转换为Android官方支持的TFLite格式即可:
- 训练完模型后执行以下代码导出:
import tensorflow as tf # 替换为你自己训练好的Keras模型变量名 trained_model = your_model # 先保存为SavedModel标准格式 trained_model.save("cat_dog_classifier") # 转换为TFLite格式 converter = tf.lite.TFLiteConverter.from_saved_model("cat_dog_classifier") # 如果用到自定义层、特殊算子,加下面这行配置 # converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS, tf.lite.OpsSet.SELECT_TF_OPS] tflite_model = converter.convert() # 保存为文件后下载到本地 with open("cat_dog_classifier.tflite", "wb") as f: f.write(tflite_model)
步骤2:导入模型到Android Studio项目
- 打开目标Android项目,右键
app模块 → New → Other → TensorFlow Lite Model,选择刚才下载的tflite文件,勾选自动导入依赖,确认后Android Studio会自动将模型放入assets目录,同时生成对应的调用类。 - 手动导入的话,把
tflite文件放到app/src/main/assets目录,再在build.gradle(Module级别)添加配置:
android { // 其他原有配置 aaptOptions { noCompress "tflite" } } dependencies { // 其他原有依赖 implementation 'org.tensorflow:tensorflow-lite:2.15.0' // 用到自定义算子时再加下面这行 // implementation 'org.tensorflow:tensorflow-lite-select-tf-ops:2.15.0' }
步骤3:实现你需要的if-else判断逻辑
只需要保证输入图片的预处理逻辑和你训练模型时的预处理完全一致(尺寸、归一化规则等),就可以直接调用模型推理:
// 初始化模型,假设自动生成的模型类名为CatDogClassifier val model = CatDogClassifier.newInstance(context) // 替换为你训练时的模型输入尺寸,示例为224*224的RGB图 val targetSize = 224 val scaledBitmap = Bitmap.createScaledBitmap(inputBitmap, targetSize, targetSize, true) // 构造输入张量 val inputTensor = TensorBuffer.createFixedSize(intArrayOf(1, targetSize, targetSize, 3), DataType.FLOAT32) val pixelValues = FloatArray(targetSize * targetSize * 3) var index = 0 for (y in 0 until targetSize) { for (x in 0 until targetSize) { val pixel = scaledBitmap.getPixel(x, y) // 归一化规则要和训练时完全一致,示例为除以255归一化到0~1区间 pixelValues[index++] = Color.red(pixel) / 255f pixelValues[index++] = Color.green(pixel) / 255f pixelValues[index++] = Color.blue(pixel) / 255f } } inputTensor.loadArray(pixelValues) // 运行推理 val output = model.process(inputTensor).outputFeature0AsTensorBuffer // 假设你训练时正类为狗,sigmoid输出概率>0.5判断为狗,否则为猫 val isDog = output.floatArray[0] > 0.5f // 实现你要的业务逻辑 if (isDog) { // 输出检测到狗 } else { // 输出检测到猫 } // 用完释放模型资源 model.close()
如果用Java开发,逻辑完全一致,仅语法有区别。如果你的模型是用softmax输出两个类别概率,取狗对应类别的概率值判断即可。
备选方案
如果TFLite适配遇到问题,你可以选以下方案:
- 把模型部署为后端接口,Android端上传图片调用接口获取结果:适合模型体积大、需要频繁更新模型的场景,缺点是需要联网。
- 将Keras模型导出为ONNX格式,使用ONNX Runtime for Android运行,操作逻辑和TFLite部署类似。
内容的提问来源于stack exchange,提问作者user15505342
相关产品推荐
相关产品推荐

