基于MNIST与Keras的手写检测程序添加自定义手写体遇形状不匹配错误
解决MNIST模型输入形状不匹配问题
错误原因分析
报错expected shape=(None, 784), found shape=(None, 84)的核心原因有两个:
- 图像通道数不匹配:
cv2.imread()默认读取彩色图像(3通道),resize后得到(28,28,3)的张量,展平后变成(28,84),和模型要求的单通道784维输入不符。 - 输入批量维度缺失:模型接受批量输入(格式为
(批量大小, 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
相关产品推荐
相关产品推荐

