我的MNIST模型无法正确识别手写数字,问题出在哪里?
我用TensorFlow的Sequential神经网络搭建了手写数字识别模型,同时用Pygame做了一个简易绘图应用来手写整数,但模型始终无法正确识别绘制的图像。我尝试过先放大绘图区域再压缩到28×28像素,但问题依然存在,想知道具体原因是什么?
神经网络代码
import os import tensorflow as tf import cv2 import matplotlib.pyplot as plt import numpy as np import paint mnist = tf.keras.datasets.mnist (x_train, y_train), (x_test, y_test) = mnist.load_data() #x_train = tf.keras.utils.normalize(x_train, axis = 1) #x_test = tf.keras.utils.normalize(x_test, axis = 1) x_train = x_train.reshape((60000, 28, 28, 1)) / 255.0 x_test = x_test.reshape((10000, 28, 28, 1)) / 255.0 model = tf.keras.models.Sequential() model.add(tf.keras.layers.Flatten(input_shape = (28,28))) # input layer, flattened image model.add(tf.keras.layers.Dense(256, activation = "relu")) # hidden model.add(tf.keras.layers.Dropout(0.5)) model.add(tf.keras.layers.Dense(128, activation = "relu")) # hidden model.add(tf.keras.layers.Dropout(0.5)) model.add(tf.keras.layers.Dense(10, activation = "softmax")) # output model.compile(optimizer = "adam", loss = "sparse_categorical_crossentropy", metrics = ["accuracy"]) model.fit(x_train, y_train, epochs = 5) model.save("Aahans.network") paint.paintcall() image = cv2.imread("/Users/aahan_bagga/Documents/DataScience/digit.png")[:,:,0] image = np.invert(np.array([image])) #black on white model.predict(image) print(f"This digit is probably {np.argmax(model.predict(image))}") #if image is not None: #print("Image loaded successfully.") #print("Image shape:", image.shape) #white_background = np.ones_like(image) * 255 # Overlay the image on the white canvas #result = cv2.addWeighted(white_background, 1, image, 0.5, 0) # Display the resulting image #cv2.imshow("Image with White Background", result) #cv2.imshow("image", image) #cv2.waitKey(0) #cv2.destroyAllWindows() #print(f"This digit is probably {np.argmax(model.predict(image))}") #else: #print("Error: Image not loaded.")
Pygame绘图应用代码
import pygame import cv2 import numpy as np def paintcall(): # Initialize Pygame pygame.init() # Constants WIDTH, HEIGHT = 512, 512 # Updated size to 28 by 28 pixels BG_COLOR = (255, 255, 255) DRAW_COLOR = (0, 0, 0) DRAW_SIZE = 30 # Adjusted size for better visibility on a small screen # Create the screen screen = pygame.display.set_mode((WIDTH, HEIGHT)) pygame.display.flip() pygame.display.set_caption("Paint App") # Create a surface to draw on drawing_surface = pygame.Surface((WIDTH, HEIGHT), pygame.SRCALPHA) # Main loop drawing = False clock = pygame.time.Clock() a = True while a: for event in pygame.event.get(): if event.type == pygame.QUIT: small_img = (cv2.cvtColor(cv2.resize(np.flipud(np.rot90(pygame.surfarray.array3d(screen))), (28, 28)), cv2.COLOR_RGB2GRAY) / 255.0) small_img = small_img.reshape(28, 28, 1) small_img_uint8 = (small_img * 255).astype(np.uint8) # Save the image using OpenCV cv2.imwrite("digit.png", small_img_uint8) #pygame.image.save(drawing_surface, "digit.png") print("Done") a = False elif event.type == pygame.MOUSEBUTTONDOWN: drawing = True elif event.type == pygame.MOUSEBUTTONUP: drawing = False elif event.type == pygame.MOUSEMOTION and drawing: pygame.draw.circle(drawing_surface, DRAW_COLOR, pygame.mouse.get_pos(), DRAW_SIZE) # Update the display screen.fill(BG_COLOR) screen.blit(drawing_surface, (0, 0)) pygame.display.flip() # Cap the frame rate clock.tick(60) #print("Finished Drawing")
可能的原因及解决办法
训练与预测的图像归一化不一致:训练时你把MNIST数据除以255缩放到0-1范围,但预测时读取图像后只做了
np.invert,没有除以255。MNIST数据是黑底白字(0为黑,255为白),你绘制的是白底黑字,反转后黑变成255、白变成0,这时候需要再除以255,让输入和训练数据范围一致:image = np.invert(np.array([image])) / 255.0图像方向/坐标系不匹配:Pygame的图像坐标系和OpenCV、MNIST的坐标系不同,你用
np.flipud(np.rot90(...))的变换可能不对,导致最终的28×28图像是旋转或翻转的,和MNIST中的数字方向不一致。可以保存变换后的图像查看,调整旋转/翻转的次数,比如改成np.rot90(..., k=3)或者去掉np.flipud,确保数字方向正确。训练epoch不足:只训练5个epoch,加上Dropout正则化,模型可能没有充分学习到手写数字的特征。可以把epoch增加到10-20,或者调整Dropout比例(比如改成0.2-0.3),提升模型拟合能力。
绘制的数字分布与MNIST不符:MNIST的数字是居中且占据大部分画布的,而你在512×512画布上绘制的数字,缩放到28×28后可能位置偏移、笔画粗细不合适。可以调整绘图时的DRAW_SIZE,或者在缩放前先把绘制的数字居中裁剪,确保和MNIST数据的分布一致。
输入维度匹配问题:模型训练时输入是(28,28,1)的单通道灰度图,预测时你读取的image是(1,28,28),虽然Flatten层可以处理,但最好统一维度,在预测时reshape成(1,28,28,1),保证和训练数据维度一致:
image = np.invert(np.array([image])).reshape(1,28,28,1) / 255.0
内容的提问来源于stack exchange,提问作者JiffyTec

