Teachable Machine导出TFLite模型推理维度不匹配及准确率为0问题
问题解决:TensorFlow Lite推理维度不匹配+准确率为0
1. 维度不匹配错误的根本原因
Google Teachable Machine导出的TFLite模型,对于灰度图像输入,期望的张量形状是(1, 96, 96, 1)(批量大小、高、宽、通道数),但用PIL打开灰度图后得到的是(96,96)的2D数组,直接赋值会导致维度不匹配。你用np.expand_dims(image, axis=2)解决了维度问题,但准确率为0是因为图像预处理没有和训练时的格式对齐。
2. 准确率为0的核心问题:预处理不匹配
Teachable Machine训练时会自动对输入图像做以下关键预处理:
- 将像素值归一化到
[0, 1](图像分类模型的默认预处理规则) - 确保输入数据类型与模型输入张量的类型一致(如uint8或float32)
你的代码缺少这些步骤,导致输入数据的分布和训练时完全不一致,模型无法识别有效特征,因此输出全0准确率。
3. 完整修复步骤
步骤1:修正图像加载与预处理逻辑
替换原代码中加载图像的部分,添加归一化、维度扩展和类型对齐操作:
# Load an image to be classified. image = Image.open(data_folder + "inputGray.png") # 强制转为灰度图(避免意外加载为RGB格式) image = ImageOps.grayscale(image) # 确保尺寸严格匹配模型输入(即使你确认尺寸一致,此步骤可避免隐性误差) image = image.resize((width, height)) # 转换为numpy数组并归一化到[0,1](对齐Teachable Machine训练时的预处理) image_np = np.array(image) / 255.0 # 扩展维度为(96,96,1),再添加批量维度(1,96,96,1)以匹配模型输入要求 image_np = np.expand_dims(image_np, axis=-1) image_np = np.expand_dims(image_np, axis=0)
步骤2:修正set_input_tensor函数
原函数直接赋值的方式可能存在类型不匹配问题,需确保输入数据类型与模型要求一致:
def set_input_tensor(interpreter, image): tensor_index = interpreter.get_input_details()[0]['index'] input_tensor = interpreter.tensor(tensor_index)() # 强制对齐输入数据类型与模型输入张量的类型 input_tensor[:] = image.astype(interpreter.get_input_details()[0]['dtype'])
步骤3:调整classify_image函数的调用
现在传入预处理后的image_np而非原始PIL图像:
# Classify the image. time1 = time.time() label_id, prob = classify_image(interpreter, image_np) time2 = time.time()
完整修正后的代码
整合所有修改后的完整可运行代码:
from tflite_runtime.interpreter import Interpreter from PIL import Image, ImageOps import numpy as np import time def load_labels(path): # Read the labels from the text file as a Python list. with open(path, 'r') as f: return [line.strip() for i, line in enumerate(f.readlines())] def set_input_tensor(interpreter, image): tensor_index = interpreter.get_input_details()[0]['index'] input_tensor = interpreter.tensor(tensor_index)() # 匹配输入数据类型与模型要求 input_tensor[:] = image.astype(interpreter.get_input_details()[0]['dtype']) def classify_image(interpreter, image, top_k=1): set_input_tensor(interpreter, image) interpreter.invoke() output_details = interpreter.get_output_details()[0] output = np.squeeze(interpreter.get_tensor(output_details['index'])) scale, zero_point = output_details['quantization'] output = scale * (output - zero_point) ordered = np.argpartition(-output, 1) return [(i, output[i]) for i in ordered[:top_k]][0] data_folder = "/home/ben/detectClouds/" model_path = data_folder + "model.tflite" label_path = data_folder + "labels.txt" interpreter = Interpreter(model_path) print("Model Loaded Successfully.") interpreter.allocate_tensors() _, height, width, _ = interpreter.get_input_details()[0]['shape'] print("Image Shape (", width, ",", height, ")") # Load an image to be classified. image = Image.open(data_folder + "inputGray.png") # 强制转为灰度图 image = ImageOps.grayscale(image) # 确保尺寸匹配 image = image.resize((width, height)) # 归一化并扩展维度 image_np = np.array(image) / 255.0 image_np = np.expand_dims(image_np, axis=-1) image_np = np.expand_dims(image_np, axis=0) # Classify the image. time1 = time.time() label_id, prob = classify_image(interpreter, image_np) time2 = time.time() classification_time = np.round(time2-time1, 3) print("Classification Time =", classification_time, "seconds.") # Read class labels. labels = load_labels(label_path) # Return the classification label of the image. classification_label = labels[label_id] print("Image Label is :", classification_label, ", with Accuracy :", np.round(prob*100, 2), "%.")
4. 额外排查点
- 若模型为量化模型(输入类型为uint8),需调整归一化方式:从
interpreter.get_input_details()[0]['quantization']获取scale和zero_point,替换归一化步骤为image_np = (np.array(image) - zero_point) * scale。 - 验证训练图像与测试图像的一致性:用PIL打开一张训练用灰度图,打印其数组的最大值、最小值,与测试图像对比,确保色彩空间、像素值范围完全匹配。
内容的提问来源于stack exchange,提问作者UltrasoundJelly
相关产品推荐
相关产品推荐

