AI眼病检测TFLite模型端点推理结果异常求助
眼部疾病分类TFLite模型输出异常排查方案
问题概述
你基于EfficientNetB3训练的眼部疾病分类模型(转TFLite格式),无论输入何种眼部图片,始终返回cataract结果,运行无报错但分类完全不准确。
核心排查方向及解决步骤
1. 验证输入预处理与训练阶段的一致性
模型的预处理逻辑必须和训练时完全匹配,这是最常见的错误来源:
- 归一化方式:你当前使用
image_array = np.array(image) / 255.0做简单归一化,但EfficientNet系列通常使用ImageNet数据集的均值和标准差做归一化:mean = np.array([0.485, 0.456, 0.406]) std = np.array([0.229, 0.224, 0.225]) image_array = (image_array / 255.0 - mean) / std - 图像通道顺序:确认训练时是用PIL(RGB)还是OpenCV(BGR)读取图像,若训练用BGR,需将PIL读取的图像转换为BGR:
image_array = image_array[..., ::-1] # RGB转BGR - Resize插值方式:明确指定插值方法,和训练时保持一致:
image = image.resize((224, 224), Image.Resampling.BILINEAR) # 或训练时用的其他插值
2. 检查模型输入输出的匹配性
添加代码打印模型的输入输出细节,确认预处理后的数组符合要求:
input_details = interpreter.get_input_details() output_details = interpreter.get_output_details() print("输入参数:", input_details) print("输出参数:", output_details)
重点确认:
- 输入的shape是否为
[1,224,224,3] - 输入的dtype是否为
float32 - 预处理后的数组维度、类型是否和上述要求完全匹配
3. 分析原始预测结果
不要只取argmax,打印所有类别的预测概率,判断是模型完全倾向于cataract,还是标签顺序错误:
print("原始预测概率:", predictions)
如果cataract的概率接近1,其他类别接近0,说明预处理错误或模型本身有问题;如果概率差距不大,可能是标签顺序和训练时不一致。
4. 验证模型转换与标签顺序
- 标签顺序:确认
data = ['cataract', 'diabetic', 'glaucoma', 'normal']的顺序和训练时的类别顺序完全一致,比如训练时类别顺序可能是['normal', 'cataract', ...],会导致argmax对应错误标签。 - 模型转换正确性:用原Keras模型测试同一张图片,若原模型结果正确,说明TFLite转换过程中出现问题,需重新转换(转换时确保启用完整的算子兼容)。
修改后的测试脚本
import numpy as np from PIL import Image import tensorflow.lite as tflite model_path = 'efficientnetb3-EyeDisease-96.22.tflite' interpreter = tflite.Interpreter(model_path) interpreter.allocate_tensors() # 打印输入输出细节,用于验证匹配性 input_details = interpreter.get_input_details() output_details = interpreter.get_output_details() print("输入参数:", input_details) print("输出参数:", output_details) image_path='1145_right.jpeg' def preprocess_image(image_path): # 匹配EfficientNet训练时的预处理逻辑 image = Image.open(image_path).convert('RGB') image = image.resize((224, 224), Image.Resampling.BILINEAR) image_array = np.array(image) # 使用ImageNet均值/std归一化 mean = np.array([0.485, 0.456, 0.406]) std = np.array([0.229, 0.224, 0.225]) image_array = (image_array / 255.0 - mean) / std image_array = np.expand_dims(image_array, axis=0) return image_array.astype(np.float32) def perform_inference(image_array): input_tensor = interpreter.tensor(input_details[0]['index']) input_tensor()[0] = image_array interpreter.invoke() output_tensor = interpreter.tensor(output_details[0]['index']) return output_tensor() if __name__ == '__main__': data = ['cataract', 'diabetic', 'glaucoma', 'normal'] image_array = preprocess_image(image_path) predictions = perform_inference(image_array) # 打印完整预测概率,便于排查 print("原始预测概率:", predictions) index = np.argmax(predictions) print('分类结果:', data[index])
内容的提问来源于stack exchange,提问作者Ahmed Bdiwy
相关产品推荐
相关产品推荐

