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()
核心问题与修复说明
- 颜色格式不匹配:MNIST是黑底白字,手绘脚本是白底黑字,需反转图像颜色。
- 缺失归一化:训练时将图像除以255归一化,预测时需同步处理。
- 模型编译配置不一致:加载模型后编译的损失函数和指标要和训练时保持一致。
- 手绘数字未居中:直接缩放50x50画布会导致数字偏移,需裁剪手绘区域后居中到28x28画布。
- 训练代码冗余错误:删除模型定义前的学习率设置代码,以及未使用的optimizer定义。
内容的提问来源于stack exchange,提问作者Catacroker
相关产品推荐
相关产品推荐

