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

TensorFlow Lite模型预测输出全为0问题求助

问题

用TensorFlow 2.10构建的分类模型在PC端运行正常,但转成TensorFlow Lite格式后,在Flutter的Android 13设备上预测输出全为0,无法定位问题原因。

Dart 数据处理与推理代码

final String column_order = List ColumnOrderData = await json.decode(column_order);
var beacons = data[0]['beacons'];
var exampleData = List<double>.filled(234, 0);

for (var i = 0; i < beacons.length; i++) {
  var beaconName = beacons[i]['beaconName'];
  double beaconRssiValue = beacons[i]['rssi'].toDouble();
  if (beaconRssiValue != 127) {
    beaconRssiValue = (beaconRssiValue * (-1)) / 105;
    int index = ColumnOrderData.indexWhere((element) => element == beaconName);
    exampleData[index] = beaconRssiValue;
  }
}

var output = List<double>.filled(4, 0);
final interpreter = await Interpreter.fromAsset('assets/model.tflite');
interpreter.run([exampleData], output.reshape([1, 4]));

TensorFlow 模型定义代码

visible = Input(shape=(234,))
hidden1 = Dense(250, activation='relu', kernel_initializer='he_normal')(visible)
drp1 = Dropout(0.5)(hidden1)
norm1 = BatchNormalization()(drp1)
hidden2 = Dense(125, activation='relu', kernel_initializer='he_normal')(norm1)
drp2 = Dropout(0.5)(hidden2)
norm2 = BatchNormalization()(drp2)
hidden3 = Dense(80, activation='relu', kernel_initializer='he_normal')(norm2)
drp3 = Dropout(0.5)(hidden3)
norm3 = BatchNormalization()(drp3)
hidden4 = Dense(40, activation='relu', kernel_initializer='he_normal')(norm3)
drp4 = Dropout(0.5)(hidden4)
norm4 = BatchNormalization()(drp4)
out_reg = Dense(4, activation='softmax')(norm4)
model = Model(inputs=visible, outputs=[out_reg])
model.compile(loss='categorical_crossentropy', optimizer=Adam(learning_rate=0.001))

TFLite 模型转换代码

tf.saved_model.save(model, 'saved_models\\')
converter = tf.lite.TFLiteConverter.from_saved_model('saved_models\\')
tflite_model = converter.convert()
with open('model.tflite', 'wb') as f:
  f.write(tflite_model)

输入输出示例

输入数据

[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.8952380952380953, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.9523809523809523, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.8952380952380953, 0.0, 0.0, 0.9238095238095239, 0.9428571428571428, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0]

输出结果

[0.0, 0.0, 0.0, 0.0]

排查与解决方向

  1. Dropout层推理状态校验:训练时Dropout会随机失活神经元,但推理阶段应自动关闭。转换前先调用model.trainable = False,确保所有层切换到推理模式再保存模型。
  2. BatchNormalization层参数确认:BatchNormalization在训练时更新均值方差,推理时需使用训练好的固定值。可在模型保存前手动调用一次model.predict(),触发层的推理模式初始化。
  3. 关闭TFLite转换优化:默认优化可能导致精度丢失,尝试强制使用float32精度:
    converter = tf.lite.TFLiteConverter.from_saved_model('saved_models\\')
    converter.optimizations = []
    converter.target_spec.supported_types = [tf.float32]
    tflite_model = converter.convert()
    
  4. 输入数据一致性检查:核对Dart代码中ColumnOrderData的加载逻辑、索引匹配规则、Rssi归一化计算,确保和PC端训练/测试时的处理完全一致。
  5. 输入输出形状匹配验证:在Dart中打印interpreter.getInputTensor(0).shape和interpreter.getOutputTensor(0).shape,确认与模型定义的(1,234)、(1,4)一致。
  6. SavedModel签名检查:使用saved_model_cli show --dir saved_models\\ --all查看模型签名,确认输入输出的形状、类型符合预期。

内容的提问来源于stack exchange,提问作者I'mStuckOnLine911

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 00:02:03