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

TensorFlow Lite图像分类:Fashion-MNIST图像存读问题排查

问题分析与解决方案

操作中的问题

  • 错误的图像保存方式:直接将numpy数组的原始二进制数据写入.unit8文件,这不是PNG、JPG这类带有格式头的标准图像文件,imread和PIL只能识别符合规范的图像格式,自然无法解析。
  • 函数参数bug:check_mnist_image_format函数定义时未接收index参数,但调用时传入了0,运行时会触发参数不匹配错误(虽当前报错为读取问题,但该函数本身存在缺陷)。
  • 文件名后缀误导:.unit8不是标准图像后缀,混淆了数据类型与文件格式的概念。

正确的图像保存与读取方法

要将Fashion-MNIST的数组保存为可识别的图像,需借助PIL将numpy数组转换为图像对象,再保存为PNG/JPG等标准格式:

import numpy as np
import matplotlib.pyplot as plt
from tensorflow.keras.datasets import fashion_mnist
import tensorflow as tf
from os.path import join
from PIL import Image
from matplotlib.pyplot import imread

# 加载数据集
(train_X, train_y), (test_X, test_y) = fashion_mnist.load_data()

# 修复图像检查函数:添加index参数
def check_mnist_image_format(image_array, index):
    print(f"Image {index} shape: {image_array.shape}")
    print(f"Image {index} data type: {image_array.dtype}")
    print(f"Image {index} min pixel value: {np.min(image_array)}")
    print(f"Image {index} max pixel value: {np.max(image_array)}")

# 检查第一张训练图像
check_mnist_image_format(train_X[0], 0)

# 正确保存图像:用PIL转换为灰度图后保存为PNG
filename = join("/home/gachaconr/tf/", 'image.png')
# Fashion-MNIST是单通道灰度图,指定mode='L'对应8位灰度格式
img = Image.fromarray(train_X[0], mode='L')
img.save(filename)
print("image saved")    

# 读取保存的图像并适配格式
image_array = imread(filename)
# 部分imread后端会将灰度图转为3通道,需转回单通道
if image_array.ndim == 3:
    image_array = image_array[:, :, 0]
check_mnist_image_format(image_array, 0)

用于TensorFlow Lite预测的完整流程

1. 构建并训练Fashion-MNIST模型

# 构建简单分类模型
model = tf.keras.Sequential([
    tf.keras.layers.Flatten(input_shape=(28, 28)),
    tf.keras.layers.Dense(128, activation='relu'),
    tf.keras.layers.Dense(10, activation='softmax')
])

# 编译模型
model.compile(optimizer='adam',
              loss='sparse_categorical_crossentropy',
              metrics=['accuracy'])

# 训练模型(训练时需归一化数据)
model.fit(train_X / 255.0, train_y, epochs=5)

2. 转换为TensorFlow Lite模型

# 保存原始Keras模型
model.save("fashion_mnist_model.h5")

# 转换为TFLite格式
converter = tf.lite.TFLiteConverter.from_keras_model(model)
tflite_model = converter.convert()

# 保存TFLite模型文件
with open("fashion_mnist_model.tflite", "wb") as f:
    f.write(tflite_model)

3. 使用保存的图像进行TFLite预测

# 加载TFLite模型
interpreter = tf.lite.Interpreter(model_path="fashion_mnist_model.tflite")
interpreter.allocate_tensors()

# 获取模型的输入、输出张量信息
input_details = interpreter.get_input_details()
output_details = interpreter.get_output_details()

# 读取并预处理图像:和训练时保持一致的归一化操作
img = Image.open(filename).convert('L')
img_array = np.array(img, dtype=np.float32) / 255.0
# 添加batch维度(模型输入要求形状为(1,28,28))
input_data = np.expand_dims(img_array, axis=0)

# 设置模型输入
interpreter.set_tensor(input_details[0]['index'], input_data)

# 执行推理
interpreter.invoke()

# 获取预测结果
output_data = interpreter.get_tensor(output_details[0]['index'])
predicted_label = np.argmax(output_data)
print(f"预测标签: {predicted_label}, 真实标签: {test_y[0]}")

内容的提问来源于stack exchange,提问作者gus

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 14:38:28