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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 03:45:00