Android TensorFlow中如何将2D传感器数据输入.pb格式预训练CNN模型
嘿,我之前刚好处理过类似的Android端传感器数据喂给CNN模型的需求,给你梳理下具体的实现步骤,应该能帮你解决问题:
核心前提:调整输入数据形状
首先得明确一个关键点:CNN模型几乎都要求输入是4D张量,格式一般是 [batch_size, height, width, channels]。你的128行6列数据,对应128个时间步(或者说样本点)、6个特征(加速度3轴+陀螺仪3轴),所以需要把这个2D数组转换成模型期望的4D形状——比如最常见的 [1, 128, 6, 1](batch_size=1表示单次输入一个样本,通道数设为1因为这是单通道的时序特征)。
具体实现步骤
1. 推荐:把.pb模型转成TensorFlow Lite格式
在Android上,TFLite是官方主推的部署方案,比直接用原生TensorFlow加载.pb模型更轻量、性能更好。你可以用Python写个小脚本把.pb转成.tflite:
import tensorflow as tf # 加载你的.pb冻结模型 graph_def = tf.compat.v1.GraphDef() with open('your_model.pb', 'rb') as f: graph_def.ParseFromString(f.read()) # 转换为TFLite,这里要替换成你模型实际的输入输出节点名 # 不知道节点名的话,可以用Netron工具打开.pb文件查看结构 converter = tf.lite.TFLiteConverter.from_graph_def( graph_def, input_arrays=['your_input_node_name'], output_arrays=['your_output_node_name'] ) tflite_model = converter.convert() # 保存转换后的模型 with open('model.tflite', 'wb') as f: f.write(tflite_model)
2. 在Android项目里集成TFLite
先在Module级别的build.gradle里添加依赖:
dependencies { // TFLite核心库,选个最新稳定版就行 implementation 'org.tensorflow:tensorflow-lite:2.15.0' // 可选:如果需要更便捷的张量处理,加这个support库 implementation 'org.tensorflow:tensorflow-lite-support:0.4.4' }
3. 把传感器的2D double数组转成模型输入张量
假设你已经拿到了double[][] sensorData(128行6列),先把它转成模型需要的float类型,再调整成4D形状:
// 第一步:把double数组转成float数组(大部分模型用float32输入) val floatData = FloatArray(128 * 6) var idx = 0 for (row in sensorData) { for (value in row) { floatData[idx++] = value.toFloat() } } // 第二步:把一维float数组包装成符合模型要求的4D张量 // 这里的形状要和你模型的输入形状完全匹配,我用的是[1,128,6,1]做例子 val inputShape = intArrayOf(1, 128, 6, 1) val inputTensor = TensorBuffer.createFixedSize(inputShape, DataType.FLOAT32) inputTensor.loadArray(floatData)
4. 加载模型并运行推理
接下来就是加载模型、喂数据、拿结果了:
// 加载assets里的tflite模型(记得把model.tflite放到assets文件夹里) val model = TensorFlowLiteModel.newInstance(context, "model.tflite") // 运行推理 val outputs = model.process(inputTensor) // 获取输出结果,outputFeature0是默认命名,如果你模型有多个输出要对应调整 val result = outputs.outputFeature0.floatArray // 用完模型一定要关闭,避免内存泄漏 model.close()
如果你非要直接用.pb模型(不推荐)
Android原生TensorFlow库现在已经不主推了,包体积大、性能差,但如果一定要用的话,步骤大概是:
- 添加原生TensorFlow依赖
- 加载.pb的GraphDef,创建Session
- 把double数组转成float张量,绑定到输入节点
- 运行Session获取输出
但真心不建议这么做,TFLite在移动端的优势太明显了。
几个要踩的坑提醒
- 输入形状必须完全匹配:如果你的模型输入是
[None,6,128,1](把特征数放前面、时间步放后面),那你就得调整数组的维度顺序,别搞反了。 - 数据类型要对:大部分模型用float32输入,别直接喂double数组,会报错。
- 节点名称不能错:一定要用Netron工具确认模型的输入输出节点名,不然加载模型的时候会找不到节点。
内容的提问来源于stack exchange,提问作者EagleJ
相关产品推荐
相关产品推荐

