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

基于Keras的自定义数据集数独数字识别模型准确率过低问题咨询

数独数字识别模型准确率异常问题排查

问题描述

正在为数独程序开发Python识别功能,目标是识别28×28的数字图片。最初参考《Python深度学习(第二版)》基于Keras框架用MNIST数据集训练模型,官方准确率可达97%~99%,但自有图片预测准确率远低于预期。后续更换为自行搭建的高清晰度专属数字数据集后,准确率仍然达不到90%,需要排查异常原因。

现有代码

模型构建代码

from tensorflow import keras
from tensorflow.keras import layers
from tensorflow.keras.preprocessing import image_dataset_from_directory
import matplotlib.pyplot as plt

def get_mnist_model_3():
    inputs = keras.Input(shape=(28, 28, 1))
    x = layers.Conv2D(filters=32, kernel_size=3, activation="relu")(inputs)
    x = layers.MaxPooling2D(pool_size=2)(x)
    x = layers.Conv2D(filters=64, kernel_size=3, activation="relu")(x)
    x = layers.MaxPooling2D(pool_size=2)(x)
    x = layers.Conv2D(filters=128, kernel_size=3, activation="relu")(x)
    x = layers.Flatten()(x)
    outputs = layers.Dense(10, activation="softmax")(x)
    model = keras.Model(inputs=inputs, outputs=outputs)
    return model

model_3 = get_mnist_model_3()
model_3.compile(optimizer="rmsprop",
    loss="sparse_categorical_crossentropy",
    metrics=["accuracy"])

callbacks_list_3 = [
    keras.callbacks.EarlyStopping(
        monitor="val_loss",
        min_delta=0,
        patience=1,),
    keras.callbacks.ModelCheckpoint(
        filepath="checkpoint_3.keras",
        monitor="val_loss",
        save_best_only=True,)]

dataset_3 = image_dataset_from_directory(
    "/home/pi/Documents/AI_Keras/sdk_numbers",
    image_size=(28, 28),
    color_mode="grayscale",
    batch_size=1)

#Total 1391 batches of size 1
split1_qty = 973 #70%
split2_qty = 347 #25%
train_dataset_3 = dataset_3.take(split1_qty)
temp_dataset = dataset_3.skip(split1_qty)
validation_dataset_3 = temp_dataset.take(split2_qty)
test_dataset_3 = temp_dataset.skip(split2_qty)

history_3 = model_3.fit(train_dataset_3,
          epochs=10,
          callbacks=callbacks_list_3,
          validation_data=(validation_dataset_3))

test_metrics_3 = model_3.evaluate(test_dataset_3)
print("test_acc  _3: {0:.1%}  /#/  test_loss _3: {1:.1%}".format(test_metrics_3[1], test_metrics_3[0]))

history_dict_3 = history_3.history
loss_values_3 = history_dict_3["loss"]
val_loss_values_3 = history_dict_3["val_loss"]
acc_3 = history_dict_3["accuracy"]
val_acc_3 = history_dict_3["val_accuracy"]
epochs_3 = range(1, len(loss_values_3) + 1)

plt.plot(epochs_3, loss_values_3, "go", label="Training loss_3")
plt.plot(epochs_3, val_loss_values_3, "g", label="Validation loss_3")
plt.title("Training and validation loss")
plt.xlabel("Epochs")
plt.ylabel("Loss")
plt.legend()
plt.show()

plt.clf()
plt.plot(epochs_3, acc_3, "go", label="Training acc_3")
plt.plot(epochs_3, val_acc_3, "g", label="Validation acc_3")
plt.title("Training and validation accuracy")
plt.xlabel("Epochs")
plt.ylabel("Accuracy")
plt.legend()
plt.show()

model_3.save('/home/pi/Documents/AI_Keras/model_3.keras')
print('Done')

模型推理代码

import random
from tensorflow import keras
from tensorflow.keras.preprocessing import image_dataset_from_directory
import pathlib
import numpy as np
import matplotlib.pyplot as plt
import cv2 as cv

model_3 = keras.models.load_model('/home/pi/Documents/AI_Keras/model_3.keras')

dataset = image_dataset_from_directory(
    "/home/pi/Documents/AI_Keras/Numbers",
    image_size=(28, 28),
    color_mode="grayscale",
    batch_size=10)

data_dir = '/home/pi/Documents/AI_Keras/Numbers'
data_dir = pathlib.Path(data_dir)
rest = list(data_dir.glob('rest/*'))

print('Target files :')
a = []
for i in range (9):
    aux = random.randint(1,len(rest))
    a.append(aux)

plt.figure(figsize=(10, 10))
for i in range (9):
    print(rest[a[i]])
    image = cv.imread(str(rest[a[i]]))
    image = cv.cvtColor(image, cv.COLOR_BGR2GRAY)
    image1 = np.array(image.reshape(28 * 28).astype("float32") / 255)
    image1 = np.expand_dims(image1, axis=0)
    image2 = np.expand_dims(image, axis=0)
    predictions_3 = model_3.predict(image2)
    pred= str(predictions_3.argmax())
    ax = plt.subplot(3, 3, i + 1)
    plt.axis("off")
    plt.title(pred)
    plt.imshow(image)

plt.show()
print('Done')

问题原因与修复方案

  • 核心问题:推理阶段预处理和训练阶段不匹配
    训练时image_dataset_from_directory会自动将图像像素值归一化到01区间,推理时你生成了归一化的`image1`变量但未使用,反而将0255区间的原始灰度数据image2输入模型,数据分布完全错位直接导致预测失效。同时模型要求输入维度为(批量数, 28, 28, 1),你生成的输入维度不符合要求。
    修复推理代码的预测部分:
    # 替换原image2生成和predict逻辑
    # 像素归一化到0~1和训练逻辑对齐
    image_norm = image.astype("float32") / 255
    # 扩展维度匹配模型输入要求:(1, 28, 28, 1)
    model_input = np.expand_dims(np.expand_dims(image_norm, axis=-1), axis=0)
    predictions_3 = model_3.predict(model_input)
    
  • 早停策略过于激进
    你设置的patience=1意味着验证损失只要一次升高就停止训练,模型很容易还没收敛就提前终止,建议调整为patience=3~5,给模型足够的迭代空间。
  • 数据集拆分逻辑不严谨
    没有固定随机种子的情况下直接用take和skip拆分数据集,可能出现训练、验证、测试集分布不一致的问题。建议在调用image_dataset_from_directory时添加seed参数固定随机状态,或者直接使用内置的validation_split参数完成数据集拆分。
  • 训练批次过小
    batch_size=1会导致训练过程梯度波动极大,模型收敛不稳定,建议调整为8~32的常规批次大小。
  • 检查颜色翻转匹配
    确认你的数据集是黑底白字还是白底黑字,如果和训练数据的前景背景颜色相反,需要在预处理阶段做对应翻转,保证数据分布一致。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 02:15:07