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

CNN无法识别手绘数字求助:模型仅适配MNIST数据集图像

手绘数字识别CNN模型故障排查

我搭建的CNN模型能正常识别MNIST数据集和测试图像,但无法正确识别鼠标手绘的数字。怀疑问题出在手绘图像的缩放环节,尝试过调整画布尺寸和模型学习率,问题仍未解决。

CNN模型代码

from tensorflow.keras import layers
from tensorflow.keras import models
from tensorflow.keras.datasets import mnist
from tensorflow.keras.utils import  to_categorical
from keras import backend as K

# 注:该行代码放在模型定义前会报错,需删除或移至模型编译后
# K.set_value(model.optimizer.learning_rate, 0.001)

(train_images, train_labels), (test_images,test_labels) = mnist.load_data()
train_images = train_images.reshape((60000, 28, 28, 1))
train_images = train_images.astype("float32")/255

test_images = test_images.reshape((10000, 28, 28, 1))
test_images = test_images.astype("float32")/255

train_labels = to_categorical(train_labels)
test_labels = to_categorical(test_labels)

model=models.Sequential()
model.add(layers.Conv2D(6,(5,5),activation="tanh",input_shape=(28,28,1)))
model.add(layers.MaxPooling2D(2,2))
model.add(layers.Conv2D(16,(5,5),activation="tanh"))
model.add(layers.MaxPooling2D(2,2))
model.add(layers.Conv2D(120,(4,4),activation="tanh"))
model.add(layers.Flatten())
model.add(layers.Dense(64,activation="tanh"))
model.add(layers.Dense(10,activation="softmax"))

# 注:该行定义的optimizer未被使用,可删除
# optimizer = keras.optimizers.Adam(lr=0.01)

model.compile(optimizer="rmsprop", loss="categorical_crossentropy", metrics=["accuracy"])
model.fit(train_images, train_labels, epochs=5, batch_size=64)

test_loss, test_acc = model.evaluate(test_images,test_labels)

model_json = model.to_json()
with open("model.json", "w") as json_file:
    json_file.write(model_json)
model.save_weights("model.h5")
print("Modelo Guardado!")

手绘数字预测脚本

import sys
from tkinter import *
from PIL import Image, ImageDraw
from keras.models import model_from_json
import keras.utils as image
import numpy as np

drawing_area=""
w=50
h=50
x,y=None,None
count=0
image_count=0
image_name="numero"
pil_image=Image.new("1",(w,h),"white")
draw=ImageDraw.Draw(pil_image)

# 加载模型
json_file = open('model.json', 'r')
loaded_model_json = json_file.read()
json_file.close()
loaded_model = model_from_json(loaded_model_json)
loaded_model.load_weights("model.h5")
print("Cargado modelo desde disco.")
# 注:此处编译配置需与训练时一致,或直接删除编译步骤
loaded_model.compile(optimizer="rmsprop", loss="categorical_crossentropy", metrics=["accuracy"])

def graficar(event):
    global drawing_area,x,y,count,draw
    newx, newy= event.x, event.y
    if x is None:
        x,y=newx, newy
        return
    count+=1
    sys.stdout.write("revent count %d" %count)
    drawing_area.create_line((x,y,newx,newy),width=5,smooth=True)
    draw.line((x,y,newx,newy),width=10)
    x,y=newx,newy

def graficar_finalizar(event):
    global x,y
    x,y=None, None

def salir(event):
    sys.exit()

def predecir(event):
    global pil_image, image_name, image_count
    image_count +=1
    file_name = image_name+str(image_count)+".jpg"

    # 新增:裁剪并居中手绘数字
    bbox = pil_image.getbbox()
    if bbox:
        pil_image = pil_image.crop(bbox)
        scale = 20 / max(pil_image.size)
        new_size = (int(pil_image.width * scale), int(pil_image.height * scale))
        pil_image = pil_image.resize(new_size, Image.Resampling.LANCZOS)
        new_img = Image.new("1", (28,28), "white")
        paste_x = (28 - pil_image.width) // 2
        paste_y = (28 - pil_image.height) // 2
        new_img.paste(pil_image, (paste_x, paste_y))
        pil_image = new_img
    else:
        print("请先绘制数字")
        return

    pil_image.save(file_name)
    img = image.load_img(file_name,color_mode="grayscale")
    img = image.img_to_array(img)

    # 新增:反转颜色(匹配MNIST黑底白字格式)+ 归一化
    img = 255 - img
    img = img.astype("float32") / 255

    img = np.expand_dims(img, axis=0)
    classes = loaded_model.predict(img)
    print(np.argmax(classes))

def limpiar(event):
    global drawing_area, pil_image, draw
    drawing_area.delete("all")
    pil_image=Image.new("1",(w,h),"white")
    draw=ImageDraw.Draw(pil_image)

def main():
    global drawing_area
    win=Tk()
    win.title("Lienzo")
    drawing_area=Canvas(win,width=w,height=h,bg="white")
    drawing_area.bind("<B1-Motion>",graficar)
    drawing_area.bind("<ButtonRelease-1>",graficar_finalizar)
    drawing_area.pack()

    b1=Button(win,text="Predecir",bg="white")
    b1.pack()
    b1.bind("<Button-1>",predecir)

    b2=Button(win,text="Limpiar",bg="white")
    b2.pack()
    b2.bind("<Button-1>",limpiar)

    b3 = Button(win, text="Cerrar", bg="white")
    b3.pack()
    b3.bind("<Button-1>", salir)

    win.mainloop()

if __name__=="__main__":
    main()

核心问题与修复说明

  1. 颜色格式不匹配:MNIST是黑底白字,手绘脚本是白底黑字,需反转图像颜色。
  2. 缺失归一化:训练时将图像除以255归一化,预测时需同步处理。
  3. 模型编译配置不一致:加载模型后编译的损失函数和指标要和训练时保持一致。
  4. 手绘数字未居中:直接缩放50x50画布会导致数字偏移,需裁剪手绘区域后居中到28x28画布。
  5. 训练代码冗余错误:删除模型定义前的学习率设置代码,以及未使用的optimizer定义。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 18:24:59