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

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

三、核心问题分析

  1. 输入维度不匹配:MNIST模型的输入格式为[batch_size, 1, 28, 28](单通道灰度图),但测试代码仅添加了batch维度,缺少通道维度,导致模型输入形状错误。
  2. 预处理未对齐MNIST风格:MNIST数据集是白底黑字的单数字图片,而门牌号图片多为黑底白字,且存在背景干扰,直接缩放会丢失数字特征。
  3. 模型适配场景差异:训练好的模型仅针对单数字识别,门牌号是多数字组合,需先分割再识别。

四、解决方案

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 19:44:56