如何向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
相关产品推荐
相关产品推荐

