如何在Android设备上无需TF Lite与服务器部署TensorFlow模型
在离线Android设备加载TensorFlow SavedModel(支持5D张量)的可行方案
Java 实现方案
步骤1:准备离线依赖
由于设备无网络,需提前下载TensorFlow原生Java库的离线包(org.tensorflow:tensorflow:2.x.x,选择适配Android的版本),将其放入项目的libs目录,在build.gradle中添加本地依赖:
dependencies { implementation files('libs/tensorflow-2.x.x.jar') // 若为aar包,使用:implementation(name: 'tensorflow-android-2.x.x', ext: 'aar') }
同时配置aaptOptions避免压缩模型文件:
android { aaptOptions { noCompress 'pb', 'pbtxt', 'savedmodel' } }
步骤2:复制SavedModel到本地目录
SavedModelBundle需要访问本地文件路径,无法直接读取assets中的压缩资源,需先将assets中的模型目录复制到应用私有目录:
import org.tensorflow.SavedModelBundle; import java.io.*; import android.content.Context; import android.content.res.AssetManager; private SavedModelBundle loadModel(Context context) throws IOException { File modelDir = new File(context.getFilesDir(), "target_saved_model"); if (!modelDir.exists()) { copyAssetDirectory(context, "saved_model", modelDir.getAbsolutePath()); } return SavedModelBundle.load(modelDir.getAbsolutePath(), "serve"); } // 递归复制assets目录内容到本地 private void copyAssetDirectory(Context context, String assetDir, String destDir) throws IOException { AssetManager am = context.getAssets(); String[] files = am.list(assetDir); File dest = new File(destDir); if (!dest.exists()) dest.mkdirs(); for (String file : files) { String srcPath = assetDir + "/" + file; String destPath = destDir + "/" + file; if (am.list(srcPath).length > 0) { copyAssetDirectory(context, srcPath, destPath); } else { InputStream in = am.open(srcPath); OutputStream out = new FileOutputStream(destPath); byte[] buffer = new byte[1024]; int read; while ((read = in.read(buffer)) != -1) { out.write(buffer, 0, read); } in.close(); out.close(); } } }
步骤3:执行推理(支持5D张量)
使用TensorFlow Java API创建5D张量并运行推理,注意用try-with-resources管理资源避免内存泄漏:
import org.tensorflow.Tensor; import org.tensorflow.Tensors; private float[] runInference(SavedModelBundle model, float[][][][][] inputData) { try (Tensor<float[][][][][]> inputTensor = Tensors.create(inputData)) { Tensor<?> output = model.session().runner() .feed("your_input_tensor_name", inputTensor) // 替换为模型实际输入节点名 .fetch("your_output_tensor_name") // 替换为模型实际输出节点名 .run() .get(0); // 根据模型输出形状初始化数组 float[][][][][] outputArray = new float[1][][][][]; // 示例形状,需匹配模型输出 output.copyTo(outputArray); // 提取结果,根据实际形状调整 return outputArray[0][0][0][0]; } }
关键注意事项
- 模型加载需在后台线程执行,避免ANR。
- 确保依赖库适配设备ABI(arm64-v8a/armeabi-v7a),下载对应架构的离线包。
- 用完
SavedModelBundle和Tensor后必须关闭,优先使用try-with-resources。
C++ 实现方案
步骤1:编译TensorFlow C++ Android库
从TensorFlow源码编译适配Android的C++库:
- 下载TensorFlow源码,配置Android NDK/SDK路径。
- 使用bazel编译指定架构的库:
# 编译arm64-v8a版本 bazel build --config=android_arm64 //tensorflow:libtensorflow.so bazel build --config=android_arm64 //tensorflow:libtensorflow_cc.so
将编译得到的.so文件放入项目src/main/jniLibs/[abi]/目录,头文件放入src/main/include/。
步骤2:配置CMakeLists.txt
cmake_minimum_required(VERSION 3.10.2) project("tf_inference") add_library( native-lib SHARED native-lib.cpp) # 引入TensorFlow库 add_library(tensorflow_cc SHARED IMPORTED) set_target_properties(tensorflow_cc PROPERTIES IMPORTED_LOCATION ${CMAKE_SOURCE_DIR}/src/main/jniLibs/${ANDROID_ABI}/libtensorflow_cc.so) add_library(tensorflow SHARED IMPORTED) set_target_properties(tensorflow PROPERTIES IMPORTED_LOCATION ${CMAKE_SOURCE_DIR}/src/main/jniLibs/${ANDROID_ABI}/libtensorflow.so) include_directories(src/main/include) find_library( log-lib log) target_link_libraries( native-lib tensorflow_cc tensorflow ${log-lib})
步骤3:C++推理代码(JNI调用)
#include <tensorflow/core/public/session.h> #include <tensorflow/core/platform/env.h> #include <tensorflow/core/framework/tensor.h> #include <android/log.h> #include <jni.h> using namespace tensorflow; extern "C" JNIEXPORT jfloatArray JNICALL Java_com_your_package_MainActivity_runTFInference( JNIEnv* env, jobject thiz, jstring model_path, jfloatArray input) { // 创建Session Session* session; Status status = NewSession(SessionOptions(), &session); if (!status.ok()) { __android_log_print(ANDROID_LOG_ERROR, "TFInfer", "Session create failed: %s", status.ToString().c_str()); return nullptr; } // 加载SavedModel const char* model_dir = env->GetStringUTFChars(model_path, nullptr); status = session->CreateSavedModel(model_dir, {"serve"}, nullptr); env->ReleaseStringUTFChars(model_path, model_dir); if (!status.ok()) { __android_log_print(ANDROID_LOG_ERROR, "TFInfer", "Load model failed: %s", status.ToString().c_str()); session->Close(); return nullptr; } // 处理5D输入张量 jfloat* input_data = env->GetFloatArrayElements(input, nullptr); int input_len = env->GetArrayLength(input); // 替换为你的模型输入形状 Tensor input_tensor(DT_FLOAT, TensorShape({1, 64, 64, 32, 32})); auto input_mapped = input_tensor.tensor<float, 5>(); // 按维度填充数据,需匹配张量形状 int idx = 0; for (int d0 = 0; d0 < input_tensor.dim_size(0); d0++) { for (int d1 = 0; d1 < input_tensor.dim_size(1); d1++) { for (int d2 = 0; d2 < input_tensor.dim_size(2); d2++) { for (int d3 = 0; d3 < input_tensor.dim_size(3); d3++) { for (int d4 = 0; d4 < input_tensor.dim_size(4); d4++) { input_mapped(d0, d1, d2, d3, d4) = input_data[idx++]; } } } } } env->ReleaseFloatArrayElements(input, input_data, 0); // 运行推理 std::vector<Tensor> outputs; status = session->Run( {{"your_input_name", input_tensor}}, {"your_output_name"}, {}, &outputs); if (!status.ok()) { __android_log_print(ANDROID_LOG_ERROR, "TFInfer", "Inference failed: %s", status.ToString().c_str()); session->Close(); return nullptr; } // 转换输出为Java数组 Tensor& output_tensor = outputs[0]; auto output_mapped = output_tensor.tensor<float, 5>(); int output_len = output_tensor.NumElements(); jfloatArray result = env->NewFloatArray(output_len); jfloat* result_data = env->GetFloatArrayElements(result, nullptr); idx = 0; for (int d0 = 0; d0 < output_tensor.dim_size(0); d0++) { for (int d1 = 0; d1 < output_tensor.dim_size(1); d1++) { for (int d2 = 0; d2 < output_tensor.dim_size(2); d2++) { for (int d3 = 0; d3 < output_tensor.dim_size(3); d3++) { for (int d4 = 0; d4 < output_tensor.dim_size(4); d4++) { result_data[idx++] = output_mapped(d0, d1, d2, d3, d4); } } } } } env->ReleaseFloatArrayElements(result, result_data, 0); session->Close(); return result; }
关键注意事项
- 编译库时需确保NDK版本与TensorFlow要求匹配(建议使用TensorFlow文档指定的NDK版本)。
- 模型文件需提前从assets复制到应用私有目录,JNI代码无法直接访问assets路径。
- 严格管理JNI内存,避免内存泄漏或野指针问题。
内容的提问来源于stack exchange,提问作者Thibault
相关产品推荐
相关产品推荐

