PyTorch训练MNIST模型测试自定义门牌号图片识别异常求助
MNIST模型在自定义门牌号图片上识别失败的问题与解决思路
一、模型训练与测试集表现
按照GeeksforGeeks的PyTorch教程训练的MNIST模型,训练阶段及测试集表现优异:
训练日志
Epoch [1/10],Loss:0.1767,Validation Loss:15.1023,Accuracy:0.95,Validation Accuracy:0.95 Epoch [2/10],Loss:0.1112,Validation Loss:11.1439,Accuracy:0.97,Validation Accuracy:0.96 Epoch [3/10],Loss:0.0711,Validation Loss:8.6009,Accuracy:0.96,Validation Accuracy:0.98 Epoch [4/10],Loss:0.0606,Validation Loss:7.4755,Accuracy:0.97,Validation Accuracy:0.98 Epoch [5/10],Loss:0.0393,Validation Loss:6.7248,Accuracy:0.99,Validation Accuracy:0.99 Epoch [6/10],Loss:0.0311,Validation Loss:8.4266,Accuracy:1.00,Validation Accuracy:0.99 Epoch [7/10],Loss:0.0388,Validation Loss:6.4547,Accuracy:1.00,Validation Accuracy:0.99 Epoch [8/10],Loss:0.0216,Validation Loss:6.4336,Accuracy:1.00,Validation Accuracy:1.00 Epoch [9/10],Loss:0.0426,Validation Loss:6.8441,Accuracy:1.00,Validation Accuracy:0.98 Epoch [10/10],Loss:0.0167,Validation Loss:6.0449,Accuracy:0.99,Validation Accuracy:1.00
测试集结果
Test Accuracy: 99.12% precision recall f1-score support 0 0.99 1.00 0.99 980 1 0.99 1.00 0.99 1135 2 0.99 0.99 0.99 1032 3 0.99 0.99 0.99 1010 4 0.99 0.99 0.99 982 5 0.99 0.99 0.99 892 6 1.00 0.99 0.99 958 7 0.99 0.99 0.99 1028 8 0.99 0.99 0.99 974 9 0.99 0.98 0.99 1009 accuracy 0.99 10000 macro avg 0.99 0.99 0.99 10000 weighted avg 0.99 0.99 0.99 10000
二、自定义图片测试问题
使用门牌号图片测试时模型输出错误,测试代码如下:
import cv2 import numpy as np from PIL import * image = cv2.imread(r"C:\Users\Sanmitha\Documents\first.jpg",0) image = cv2.resize(image, (28,28)) batch = torch.tensor(image / 255).unsqueeze(0).float() with torch.no_grad(): batch = batch.to(device) output = model( batch ) output = torch.argmax(output, 1) print(output)
输出结果:
tensor([5])
待识别图片:
- first.jpg:包含目标数字38/42
- second.jpg:包含目标数字110/50
三、核心问题分析
- 输入维度不匹配:MNIST模型的输入格式为
[batch_size, 1, 28, 28](单通道灰度图),但测试代码仅添加了batch维度,缺少通道维度,导致模型输入形状错误。 - 预处理未对齐MNIST风格:MNIST数据集是白底黑字的单数字图片,而门牌号图片多为黑底白字,且存在背景干扰,直接缩放会丢失数字特征。
- 模型适配场景差异:训练好的模型仅针对单数字识别,门牌号是多数字组合,需先分割再识别。
四、解决方案
1. 修正输入维度与预处理
调整代码,确保输入格式与训练时一致,同时对齐图片风格:
import cv2 import torch # 读取单通道灰度图 image = cv2.imread(r"C:\Users\Sanmitha\Documents\first.jpg", 0) # 反转颜色,转为白底黑字(匹配MNIST风格) image = cv2.bitwise_not(image) # 二值化降噪,增强数字边缘 _, image = cv2.threshold(image, 100, 255, cv2.THRESH_BINARY) # 缩放到28x28 image = cv2.resize(image, (28, 28)) # 调整维度:添加通道维度和batch维度,最终形状为[1,1,28,28] batch = torch.tensor(image / 255).unsqueeze(0).unsqueeze(0).float() with torch.no_grad(): batch = batch.to(device) output = model(batch) pred = torch.argmax(output, 1).item() print(f"预测结果:{pred}")
2. 多数字分割与识别
针对门牌号的多数字特性,先分割每个数字再逐一识别:
import cv2 import torch # 读取并预处理图片 image = cv2.imread(r"C:\Users\Sanmitha\Documents\first.jpg", 0) image = cv2.bitwise_not(image) _, image = cv2.threshold(image, 100, 255, cv2.THRESH_BINARY) # 检测数字轮廓 contours, _ = cv2.findContours(image, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) # 按x轴排序轮廓(保证数字顺序正确) contours = sorted(contours, key=lambda x: cv2.boundingRect(x)[0]) predicted_digits = [] for cnt in contours: x, y, w, h = cv2.boundingRect(cnt) # 过滤过小的噪声轮廓 if w > 15 and h > 20: # 提取数字区域 digit_roi = image[y:y+h, x:x+w] # 缩放并补边,保证数字居中(匹配MNIST的数字位置) digit_roi = cv2.resize(digit_roi, (20,20)) padded_digit = cv2.copyMakeBorder(digit_roi, 4,4,4,4, cv2.BORDER_CONSTANT, value=0) # 转为模型输入格式 batch = torch.tensor(padded_digit /255).unsqueeze(0).unsqueeze(0).float().to(device) # 推理 with torch.no_grad(): output = model(batch) pred = torch.argmax(output,1).item() predicted_digits.append(str(pred)) print(f"门牌号预测结果:{''.join(predicted_digits)}")
3. 可选:模型微调
如果预处理后识别效果仍不佳,可以收集门牌号数字图片,对现有MNIST模型进行微调,让模型适配真实场景的数字风格。
内容的提问来源于stack exchange,提问作者Summer Project
相关产品推荐
相关产品推荐

