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

基于MNIST与Keras的手写检测程序添加自定义手写体遇形状不匹配错误

解决MNIST模型输入形状不匹配问题

错误原因分析

报错expected shape=(None, 784), found shape=(None, 84)的核心原因有两个:

  1. 图像通道数不匹配:cv2.imread()默认读取彩色图像(3通道),resize后得到(28,28,3)的张量,展平后变成(28,84),和模型要求的单通道784维输入不符。
  2. 输入批量维度缺失:模型接受批量输入(格式为(批量大小, 784)),单张图像处理后未添加批量维度,导致形状不兼容。

具体修复步骤

1. 将彩色图像转为单通道灰度图

在读取图像后添加灰度转换步骤,对齐MNIST数据集的单通道格式:

img = cv2.imread(path)
img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)  # 转为灰度图

2. 修正图像reshape方式

单张图像处理完成后,添加批量维度(格式为(1, 784)),匹配模型输入要求:

image = image.reshape(1, -1)  # 从(28,28)转为(1,784)

3. 修复标题格式化错误

原代码标题语法错误,改用字符串格式化方式:

plt.title(f"Predicted: {y_sample_pred_class}", fontsize=16)

修改后的完整代码

#Libraries to import:
import numpy as np
import matplotlib.pyplot as plt
import keras
from keras.models import Sequential
from keras.layers import Dense, Dropout
from keras.datasets import mnist
import tensorflow as tf
from tensorflow import keras
import cv2
np.random.seed(0)

#Converting input image
path = r'theImage_1.png'
img = cv2.imread(path)
img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)  # 转为灰度图,匹配MNIST格式
twentyEight = cv2.resize(img, (28, 28), interpolation=cv2.INTER_LINEAR)
image = cv2.bitwise_not(twentyEight)

#Downloading data
(x_train, y_train), (x_test, y_test) = mnist.load_data()

#Categorizing data:
y_train = keras.utils.to_categorical(y_train, 10)
y_test = keras.utils.to_categorical(y_test, 10)

#Normalizing
x_train = x_train/255
x_test = x_test/255
image = image/255

#Reshaping
x_train = x_train.reshape(x_train.shape[0], -1)
x_test = x_test.reshape(x_test.shape[0], -1)
image = image.reshape(1, -1)  # 添加批量维度,转为(1,784)

#The neural network
model = Sequential()
model.add(Dense(units=128, input_shape=(784,), activation='relu'))
model.add(Dense(units=128, activation='relu'))
model.add(Dropout(0.25))
model.add(Dense(units=10, activation='softmax'))
model.compile(loss='categorical_crossentropy', optimizer='adam', metrics=['accuracy'])

#Training
model.fit(x=x_train, y=y_train, batch_size=512, epochs=10)

#Example
y_pred = model.predict(image)
y_pred_classes = np.argmax(y_pred, axis=1)
y_sample_pred_class = y_pred_classes[0]
plt.title(f"Predicted: {y_sample_pred_class}", fontsize=16)  # 修正标题格式
plt.imshow(image.reshape(28, 28), cmap='gray')
plt.show()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 23:20:28