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

在Android中运行TensorFlow模型:鸢尾花分类器部署问题

在Android中部署Keras鸢尾花分类MLP模型的完整指南

我看你正在尝试把自己用Keras训练的鸢尾花分类MLP模型部署到Android应用里,正好我在这类移动ML部署上有不少经验,给你梳理一下完整的流程和需要注意的坑,先把你没写完的模型代码补全并修正,再一步步讲Android端的实现:

一、补全并修正你的Keras模型代码

你提供的代码最后一行有拼写错误(activatio应该是activation),而且缺少模型编译、训练和保存的关键步骤,完整的可运行代码如下:

import numpy as np
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
from tensorflow.keras.layers import Input, Dense
from tensorflow.keras.models import Model

# 加载鸢尾花数据集
data = load_iris()
x = np.array(data.data)
y = np.array(data.target)

# 划分训练集和测试集,测试集占25%
x_train, x_test, y_train, y_test = train_test_split(x, y, test_size=0.25)

# 构建MLP网络结构
inputs = Input(shape=(4,))
x = Dense(10, activation="relu", name="input_layer")(inputs)
x = Dense(10, activation="relu")(x)
x = Dense(15, activation="relu")(x)
# 输出层用softmax,对应3类鸢尾花的概率分布
outputs = Dense(3, activation="softmax", name="output_layer")(x)

model = Model(inputs=inputs, outputs=outputs)

# 编译模型:因为标签是整数形式,用sparse_categorical_crossentropy损失
model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy'])

# 训练模型:跑50轮,批量大小8,用10%的训练数据做验证
model.fit(x_train, y_train, epochs=50, batch_size=8, validation_split=0.1)

# 在测试集上评估模型性能
test_loss, test_acc = model.evaluate(x_test, y_test)
print(f"测试集准确率: {test_acc:.4f}")

# 保存为SavedModel格式(包含.pb文件的文件夹),方便后续转换
model.save("iris_classifier")

二、模型转换与Android端准备

虽然你提到把.pb文件放到assets文件夹,但更推荐将模型转换为**TensorFlow Lite (.tflite)**格式——它是谷歌针对移动设备优化的轻量级格式,加载速度更快、内存占用更低。转换代码如下:

import tensorflow as tf

# 从SavedModel文件夹加载模型
converter = tf.lite.TFLiteConverter.from_saved_model("iris_classifier")
# 可选:启用整数量化进一步缩小模型体积(如果不需要高精度的话)
# converter.optimizations = [tf.lite.Optimize.DEFAULT]
tflite_model = converter.convert()

# 保存为.tflite文件
with open("iris_classifier.tflite", "wb") as f:
    f.write(tflite_model)

把转换好的tflite文件放到Android项目的src/main/assets文件夹下(如果没有assets文件夹,右键app模块→New→Folder→Assets Folder创建)。

三、Android端实现模型加载与推理

1. 添加TensorFlow Lite依赖

在app模块的build.gradle(Module级别)的dependencies块中添加依赖:

dependencies {
    // 基础TFLite依赖
    implementation 'org.tensorflow:tensorflow-lite:2.15.0'
    // 可选:如果需要GPU加速来提升推理速度
    implementation 'org.tensorflow:tensorflow-lite-gpu:2.15.0'
}

2. 编写分类器工具类

创建一个封装模型加载和推理的工具类,方便在Activity/Fragment中调用:

import android.content.Context;
import org.tensorflow.lite.Interpreter;
import java.io.FileInputStream;
import java.io.IOException;
import java.nio.MappedByteBuffer;
import java.nio.channels.FileChannel;

public class IrisClassifier {
    private Interpreter tfliteInterpreter;

    public IrisClassifier(Context context) throws IOException {
        // 从assets加载模型文件
        MappedByteBuffer modelBuffer = loadModelFromAssets(context);
        // 初始化TFLite解释器
        tfliteInterpreter = new Interpreter(modelBuffer);
    }

    // 读取assets中的模型文件,转换为MappedByteBuffer
    private MappedByteBuffer loadModelFromAssets(Context context) throws IOException {
        FileInputStream inputStream = new FileInputStream(
                context.getAssets().openFd("iris_classifier.tflite").getFileDescriptor()
        );
        FileChannel fileChannel = inputStream.getChannel();
        long startOffset = context.getAssets().openFd("iris_classifier.tflite").getStartOffset();
        long fileLength = context.getAssets().openFd("iris_classifier.tflite").getDeclaredLength();
        return fileChannel.map(FileChannel.MapMode.READ_ONLY, startOffset, fileLength);
    }

    // 核心推理方法:输入4个鸢尾花特征,返回分类结果(0/1/2对应三种鸢尾花)
    public int classifyIris(float[] flowerFeatures) {
        // 输入张量形状:[1,4](批量大小1,特征数4)
        float[][] inputTensor = new float[1][4];
        inputTensor[0] = flowerFeatures;

        // 输出张量形状:[1,3](每个类别的概率值)
        float[][] outputTensor = new float[1][3];

        // 执行推理
        tfliteInterpreter.run(inputTensor, outputTensor);

        // 找到概率最大的类别索引
        int predictedClass = 0;
        float maxProbability = outputTensor[0][0];
        for (int i = 1; i < 3; i++) {
            if (outputTensor[0][i] > maxProbability) {
                maxProbability = outputTensor[0][i];
                predictedClass = i;
            }
        }
        return predictedClass;
    }

    // 释放资源,避免内存泄漏
    public void close() {
        if (tfliteInterpreter != null) {
            tfliteInterpreter.close();
        }
    }
}

3. 在Activity中调用分类器

举个简单的调用示例,在你的主Activity中测试模型:

import android.os.Bundle;
import android.widget.TextView;
import androidx.appcompat.app.AppCompatActivity;
import java.io.IOException;

public class MainActivity extends AppCompatActivity {
    private TextView resultTextView;

    @Override
    protected void onCreate(Bundle savedInstanceState) {
        super.onCreate(savedInstanceState);
        setContentView(R.layout.activity_main);
        resultTextView = findViewById(R.id.result_text);

        try {
            // 初始化分类器
            IrisClassifier classifier = new IrisClassifier(this);
            // 示例输入:花萼长5.1,宽3.5;花瓣长1.4,宽0.2(对应Iris Setosa,类别0)
            float[] sampleFeatures = {5.1f, 3.5f, 1.4f, 0.2f};
            int predictedClass = classifier.classifyIris(sampleFeatures);

            // 映射类别到花种名称
            String flowerName = switch (predictedClass) {
                case 0 -> "Iris Setosa";
                case 1 -> "Iris Versicolor";
                case 2 -> "Iris Virginica";
                default -> "未知花种";
            };
            resultTextView.setText("预测结果:" + flowerName);

            // 记得关闭分类器释放资源
            classifier.close();
        } catch (IOException e) {
            e.printStackTrace();
            resultTextView.setText("模型加载失败");
        }
    }
}

四、常见问题排查

  • 模型加载失败:检查assets文件夹是否正确创建,模型文件名是否和代码中完全一致(Android对文件名大小写敏感)。
  • 推理结果错误:
    • 如果训练时对输入特征做了归一化/标准化,Android端推理时必须对输入做完全相同的预处理,比如训练时把特征缩放到[0,1],推理时也要把输入的特征值做同样的缩放。
    • 验证TFLite模型的推理结果和原Keras模型是否一致,避免转换过程中出现异常。
  • 性能问题:如果模型运行卡顿,可以开启TFLite的GPU加速,或者使用整数量化优化模型体积和速度。

内容的提问来源于stack exchange,提问作者merch

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 09:29:47