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

我的MNIST模型无法正确识别手写数字,问题出在哪里?

问题:手写数字识别模型无法正确识别Pygame绘制的图像

我用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 18:04:52