TFLite后训练量化模型始终输出预测概率1.0问题求助
问题排查与解决方案
核心问题分析
原Keras二分类CNN模型在测试集表现正常,但经过float16后训练量化得到的TFLite模型,无论输入何种图像都返回预测概率1.0。结合代码和张量详情来看,问题主要集中在输入数据预处理缺失和冗余的张量操作上。
1. 输入数据未做归一化(最关键原因)
训练Keras模型时,输入图像必然经过了归一化处理(比如除以255将像素值缩放到[0,1]区间,或使用特定均值/std标准化),但当前代码中仅将load_img得到的uint8格式图像直接转为float32就输入模型,像素值范围仍为[0,255],与模型训练时的输入分布完全不匹配,导致模型输出异常。
解决方案:
添加与训练代码一致的归一化步骤。例如:
- 如果训练时是除以255归一化:
test_ds = tf.keras.utils.load_img(os.path.join(path_data, "bg", img), target_size=(224, 224, 3)) # 添加除以255的归一化 test_ds = np.expand_dims(np.array(test_ds, dtype=np.float32) / 255.0, axis=0)
- 如果训练时使用了
tf.keras.applications系列模型的预处理逻辑:
test_ds = tf.keras.utils.load_img(os.path.join(path_data, "bg", img), target_size=(224, 224, 3)) test_ds = tf.keras.applications.mobilenet.preprocess_input(np.array(test_ds, dtype=np.float32)) test_ds = np.expand_dims(test_ds, axis=0)
2. 移除冗余的张量Resize操作
从in_details和out_details可以看到,模型默认输入形状就是(1,224,224,3),输出形状是(1,1),手动调用resize_tensor_input属于多余操作,反而可能破坏模型的张量分配逻辑。
解决方案:
删除以下冗余代码:
interpreter.resize_tensor_input(in_details[0]["index"], (1, 224, 224, 3)) interpreter.resize_tensor_input(out_details[0]["index"], (1, 1)) interpreter.allocate_tensors()
修改后的推理代码如下:
test_ds = tf.keras.utils.load_img(os.path.join(path_data, "bg", img), target_size=(224, 224, 3)) test_ds = np.expand_dims(np.array(test_ds, dtype=np.float32) / 255.0, axis=0) # Load interpreter interpreter = tf.lite.Interpreter(model_content=model) interpreter.allocate_tensors() in_details = interpreter.get_input_details() out_details = interpreter.get_output_details() # Inference interpreter.set_tensor(in_details[0]["index"], test_ds) interpreter.invoke() pred = interpreter.get_tensor(out_details[0]["index"])[0][0]
3. 验证量化模型转换正确性
确认float16量化是否正确应用,可通过Python API重新转换模型,确保参数配置正确:
# 加载训练好的Keras模型 keras_model = tf.keras.models.load_model("your_keras_model_path.h5") # 配置TFLite转换器 converter = tf.lite.TFLiteConverter.from_keras_model(keras_model) converter.optimizations = [tf.lite.Optimize.DEFAULT] converter.target_spec.supported_types = [tf.float16] # 转换并保存模型 tflite_model = converter.convert() with open("float16_quantized_model.tflite", "wb") as f: f.write(tflite_model)
4. 校验输入数据的维度与类型
在推理前添加校验代码,确保输入数据的形状和类型与模型要求一致:
print("Input shape:", test_ds.shape) # 应输出 (1, 224, 224, 3) print("Input dtype:", test_ds.dtype) # 应输出 float32
内容的提问来源于stack exchange,提问作者peacer
相关产品推荐
相关产品推荐

