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

基于ResNet50的MNIST模型在自定义手写数据上表现极差的求助

自定义手写数字图片预测失效的原因与解决方法

问题背景

使用PyTorch预训练ResNet50微调MNIST数据集,模型定义如下:

from torch import nn
from torchvision.models import ResNet50_Weights, resnet50

class Model(nn.Module):
  def __init__(self):
    super(Model, self).__init__()

    self.model = resnet50(weights=ResNet50_Weights.DEFAULT)

    self.model.conv1 = nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3, bias=False)
    
    num_ftrs = self.model.fc.in_features
    self.model.fc = nn.Linear(num_ftrs, 10)

  def forward(self, x):
    return self.model(x)

训练10轮后,在50000张训练图像上达到99.895%的准确率,测试代码及结果:

model.eval()

with torch.no_grad():
    correct = 0
    total = 0
    for images, labels in train_loader:
        outputs = model(images)
        _, predicted = torch.max(outputs.data, 1)
        total += labels.size(0)
        correct += (predicted == labels).sum().item()
    
    print('Accuracy of the network on the {} train images: {} %'.format(50000, 100 * correct / total))
[out]: Accuracy of the network on the 50000 train images: 99.895 %

随后用Pygame制作自定义手写数字图片,核心绘制代码:

if event.type == pg.MOUSEMOTION:
    if (drawing):
        mouse_position = pg.mouse.get_pos()
        pg.draw.circle(screen, color, mouse_position, w)
elif event.type == pg.MOUSEBUTTONUP:
    mouse_position = (0, 0)
    drawing = False
    last_pos = None
elif event.type == pg.MOUSEBUTTONDOWN:
    drawing = True

将图片转为28x28灰度图并转成张量:

image = Image.open("image.png").convert("L").resize((28,28),Image.Resampling.LANCZOS)

transform = Compose([
    PILToTensor(),
    Lambda(lambda image: image.view(-1, 1, 28, 28))
])

img_tensor = transform(image).to(torch.float)

输入模型后预测效果极差,例如输入数字2的图片,输出:

with torch.no_grad():
    outputs = model(img_tensor)
    print(outputs)
    _, predicted = torch.max(outputs.data, 1)
    print(predicted)
[out]: tensor([[ 20.6237,   0.4952, -15.5033,   8.5165,   1.0938,   2.8278,   2.0153,
           3.2825,  -6.2655,  -0.6992]])
tensor([0])

类别2的置信度为负数,模型错误预测为0。

核心原因分析

  • 像素分布完全反转:MNIST数据集的图片是白底黑字(背景灰度值255,数字灰度值0),而Pygame绘制的图片默认是黑底白字(背景0,数字255),两者像素分布完全相反,模型无法识别。
  • 缺少归一化处理:训练MNIST时通常会对数据做归一化(比如除以255缩放到[0,1],或者用MNIST的均值0.1307、标准差0.3081做标准化),但自定义图片仅转成float类型,未做任何归一化,输入数据的数值范围和训练集不匹配。
  • 图像形态差异:Pygame绘制的数字可能在位置、大小、笔触粗细上和MNIST样本差异过大,比如数字偏离中心、过于粗大/细小,导致模型无法提取到匹配的特征。

解决步骤

1. 反转灰度值,匹配MNIST像素分布

在预处理时反转图像的灰度值,将黑底白字转为白底黑字:

# 方法1:在PIL转换时处理
image = Image.open("image.png").convert("L").resize((28,28),Image.Resampling.LANCZOS)
image = Image.eval(image, lambda x: 255 - x)  # 反转灰度值

# 方法2:在张量层面处理
img_tensor = 255.0 - img_tensor

2. 添加归一化,和训练时的预处理保持一致

假设训练时的预处理包含归一化,比如:

# 训练时的transform
train_transform = Compose([
    ToTensor(),
    Normalize((0.1307,), (0.3081,))
])

那么预测时的transform必须同步添加归一化:

from torchvision.transforms import Compose, PILToTensor, Lambda, Normalize

transform = Compose([
    PILToTensor(),
    Lambda(lambda x: 255.0 - x),  # 反转灰度
    Lambda(lambda x: x / 255.0),  # 缩放到[0,1]
    Normalize((0.1307,), (0.3081,)),  # 用MNIST的均值和标准差标准化
    Lambda(lambda image: image.view(-1, 1, 28, 28))
])

img_tensor = transform(image).to(torch.float)

3. 优化自定义图片的绘制规范

  • 调整Pygame画布大小为接近28x28的比例(比如280x280,方便缩放后不失真),绘制时尽量让数字居中,大小和MNIST样本接近(MNIST数字通常占据画布的70%左右空间)。
  • 调整笔触粗细,避免数字过粗导致占满画布,或过细导致缩放后丢失细节。

4. 验证预处理后的张量一致性

可以打印训练集中一个样本的张量信息,和自定义图片预处理后的张量对比,确保数值范围、分布一致:

# 打印训练样本的信息
for images, labels in train_loader:
    print("训练样本均值:", images.mean())
    print("训练样本最大值:", images.max())
    print("训练样本最小值:", images.min())
    break

# 打印自定义图片预处理后的信息
print("自定义图片张量均值:", img_tensor.mean())
print("自定义图片张量最大值:", img_tensor.max())
print("自定义图片张量最小值:", img_tensor.min())

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 16:36:06