PyTorch高准确率手写字母模型在React+FastAPI Web应用集成时失效,持续仅预测'r'或'q'
PyTorch高准确率手写字母模型在React+FastAPI Web应用集成时失效,持续仅预测'r'或'q'
兄弟,这种训练时准确率拉满,一部署到Web就彻底拉胯的情况真的太闹心了!我帮你梳理几个最可能的原因,一个个排查应该就能解决问题:
1. 输入数据的预处理完全不一致(最常见的坑)
训练时模型对输入的要求是非常苛刻的,Web端的手写图像如果和训练集的预处理流程不一样,模型根本认不出来:
- 尺寸匹配问题:你用的是ResNet18,默认输入尺寸是224×224,但手写字母数据集(比如EMNIST)通常是28×28。你训练时是不是用了
transforms.Resize((224,224))?如果Web端没有把画布上的图像缩放到同样的尺寸,模型看到的就是“变形”的字母,肯定乱预测。 - 通道与像素值范围:你把ResNet的conv1改成了单通道输入(
nn.Conv2d(1, 64,...)),说明训练时用的是灰度图。那Web端的画布输出是不是RGB图?有没有转成单通道灰度?另外,训练时是不是把像素值归一化到了0-1(比如除以255),而Web端直接传了0-255的原始值?这会让模型的输入分布完全偏离训练时的情况。 - 图像反转问题:训练集里的手写字母是不是“黑字白底”,而你在Web画布上是“白字黑底”?反过来也一样,像素值完全反转的话,模型看到的就是完全陌生的图案。
2. 模型推理时没切到评估模式
训练时模型处于train()模式,会启用Dropout、BatchNorm的动态均值更新等机制,但推理时必须切换到eval()模式,否则输出会极度不稳定,甚至出现固定预测某几个类的情况。
检查你的FastAPI代码里,加载模型后有没有加这两行:
model.eval() with torch.no_grad(): # 关闭梯度计算,既提速又避免干扰 # 执行预测逻辑
3. 类别映射表搞反了
训练时的标签顺序和你Web端用来解析预测结果的映射表是不是对不上?比如:
- 你训练时用的EMNIST数据集,ByClass和ByMerge的类别顺序是不一样的,可能你以为索引0对应'a',实际对应了其他字母。
- 模型输出的是类别索引,你是不是把索引和字母的对应关系写错了?比如把索引17对应成'r',但实际训练时索引17是别的字母,导致所有预测都被错误解析成'r'或'q'。
4. 模型保存/加载出问题了
有没有正确保存和加载模型?比如:
- 训练时用GPU训练,加载时在CPU上运行,有没有加
map_location='cpu'?比如:
model.load_state_dict(torch.load('model.pth', map_location=torch.device('cpu')))
- 是不是保存了整个模型而不是
state_dict?正确的保存方式应该是torch.save(model.state_dict(), 'model.pth'),加载时用model.load_state_dict(...),而不是直接torch.load('model.pth')。
快速排查步骤
给你个简单的排查流程,能快速定位问题:
- 找一张训练集里你确定标签的图像,用Web端的预处理流程(比如转尺寸、转通道、归一化)处理一遍,然后输入模型看预测结果。如果和训练时的结果一致,说明模型本身没问题,问题出在Web端的用户输入预处理;如果不一致,说明模型加载或推理模式有问题。
- 打印Web端输入模型的张量信息:形状(比如是不是
[1,1,224,224])、像素值范围(比如是不是0-1)、均值方差,和训练时的输入对比,确保完全一致。 - 直接打印模型输出的原始索引,再对照训练时的类别映射表,看是不是解析成字母时出错了。
备注:内容来源于stack exchange,提问作者bahaeddine hemmem CHAREN
相关产品推荐
相关产品推荐

