You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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++库:

  1. 下载TensorFlow源码,配置Android NDK/SDK路径。
  2. 使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.21 19:45:44