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

