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

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库现在已经不主推了,包体积大、性能差,但如果一定要用的话,步骤大概是:

  1. 添加原生TensorFlow依赖
  2. 加载.pb的GraphDef,创建Session
  3. 把double数组转成float张量,绑定到输入节点
  4. 运行Session获取输出

但真心不建议这么做,TFLite在移动端的优势太明显了。

几个要踩的坑提醒
  • 输入形状必须完全匹配:如果你的模型输入是[None,6,128,1](把特征数放前面、时间步放后面),那你就得调整数组的维度顺序,别搞反了。
  • 数据类型要对:大部分模型用float32输入,别直接喂double数组,会报错。
  • 节点名称不能错:一定要用Netron工具确认模型的输入输出节点名,不然加载模型的时候会找不到节点。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 10:23:57