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

如何向TensorFlow Lite模型传入动态尺寸数组?

实现TensorFlow Lite动态尺寸输入的步骤

1. 修改Keras模型的输入层

把模型输入层里的固定长度维度改成None,就能接受任意长度的序列了。比如原来输入层写的是:

input_layer = tf.keras.layers.Input(shape=(100, 3))  # 固定死100个时间步

改成:

input_layer = tf.keras.layers.Input(shape=(None, 3))  # 不管多少个时间步都能接,每个步带X/Y/Z三个值

改完之后重新训练或加载模型,先测试下模型能不能处理不同长度的输入。

2. 转成支持动态形状的TFLite模型

转换的时候得给转换器加些配置,允许动态形状:

import tensorflow as tf

# 先加载你的Keras模型
model = tf.keras.models.load_model("your_keras_model.h5")

converter = tf.lite.TFLiteConverter.from_keras_model(model)
# 开优化(可选,但开了更好)
converter.optimizations = [tf.lite.Optimize.DEFAULT]
# 如果模型用到了TFLite内置不支持的操作,得加上这条,让它用TensorFlow原生操作
converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS, tf.lite.OpsSet.SELECT_TF_OPS]
# 指定输入形状:第一个维度是batch_size(Android里一般用1就行),第二个维度设为None表示动态
# 把"input_1"换成你模型的输入层名字,用model.input_names就能查到
converter.input_shapes = {model.input_names[0]: (1, None, 3)}

tflite_model = converter.convert()

# 保存模型
with open("model-dynamic.tflite", "wb") as f:
    f.write(tflite_model)

3. Android端处理动态输入

在Android里,每次收集到多少加速度数据,就动态设置输入张量的形状:

先加载模型

private Interpreter tflite;

private void loadModel(Context context) throws IOException {
    ByteBuffer modelBuffer = loadModelFile(context);
    tflite = new Interpreter(modelBuffer);
}

// 读模型文件的辅助方法
private ByteBuffer loadModelFile(Context context) throws IOException {
    AssetFileDescriptor fileDescriptor = context.getAssets().openFd("model-dynamic.tflite");
    FileInputStream inputStream = new FileInputStream(fileDescriptor.getFileDescriptor());
    FileChannel fileChannel = inputStream.getChannel();
    long startOffset = fileDescriptor.getStartOffset();
    long declaredLength = fileDescriptor.getDeclaredLength();
    return fileChannel.map(FileChannel.MapMode.READ_ONLY, startOffset, declaredLength);
}

动态输入+推理

// 假设accData是收集到的加速度数据列表,每个元素存的是X/Y/Z的float数组
public float[] predictActivity(List<float[]> accData) {
    int sampleCount = accData.size();
    float[] inputArray = new float[sampleCount * 3];
    
    // 把列表转成一维数组,按X/Y/Z的顺序塞进去
    int index = 0;
    for (float[] point : accData) {
        inputArray[index++] = point[0]; // X
        inputArray[index++] = point[1]; // Y
        inputArray[index++] = point[2]; // Z
    }
    
    // 动态设置输入形状:[batch_size, sampleCount, 3]
    Tensor inputTensor = Tensor.create(new long[]{1, sampleCount, 3}, DataType.FLOAT32);
    inputTensor.loadArray(inputArray);
    
    // 准备输出张量,形状跟着你的模型来,比如分类任务就设成类别数
    int numClasses = 5; // 换成你的模型实际输出的类别数
    float[] outputArray = new float[numClasses];
    Tensor outputTensor = Tensor.create(new long[]{1, numClasses}, DataType.FLOAT32);
    outputTensor.loadArray(outputArray);
    
    // 跑推理
    tflite.run(inputTensor, outputTensor);
    
    // 读结果
    outputTensor.readArray(outputArray);
    
    // 记得关张量
    inputTensor.close();
    outputTensor.close();
    
    return outputArray;
}

要注意的点

  • 模型里的层得支持动态形状,像LSTM、GRU、全局平均池化这些都行,但自己写的自定义层可能不行,得换成内置层。
  • 要是转换的时候报操作不支持的错,就确认加了tf.lite.OpsSet.SELECT_TF_OPS,还要在Android的build.gradle里加依赖:implementation 'org.tensorflow:tensorflow-lite-select-tf-ops:2.15.0',版本和你用的TensorFlow对应上。
  • 动态输入可能会影响速度,最好给输入长度设个合理范围,别搞太长的输入拖慢性能。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 16:25:25