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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 20:24:55