PyTorch实战:训练CNN手写数字识别并部署FastAPI服务

PyTorch实战:训练CNN手写数字识别并部署FastAPI服务 这次我们不谈孤立的理论公式直接把一个能跑通的神经网络项目端到端做出来。项目选型上用 PyTorch 训练一个手写数字识别模型覆盖 MNIST 数据加载、CNN 模型定义、反向传播训练、模型保存、单张图片推理、FastAPI 接口封装、批量预测脚本、显存监控这一整条链路。目标是让读者跟着代码走一遍理解神经网络从“前向计算”到“反向更新”在工程上到底是怎么落地的。这套方案的门槛不高CPU 就能完成训练只是慢一些有 NVIDIA 显卡且正确安装了 CUDA 版 PyTorch训练速度会明显提升。代码全部给出可复制版本接口部分也提供标准 HTTP 调用示例。文章会重点讲清楚三个问题网络结构每一层对应什么代码训练循环里 optimizer.zero_grad() 和 loss.backward() 到底做了什么以及模型训练完成后如何作为服务对外提供能力。适合有 Python 基础、想快速跨过神经网络代码门槛的开发者阅读。1. 神经网络实操项目核心能力速览能力项说明项目类型神经网络入门实操 完整可运行代码技术栈Python、PyTorch、Torchvision、FastAPI、Pillow数据集MNIST 手写数字数据集10 分类模型结构两层卷积网络 SimpleCNN包含卷积、池化、全连接、Dropout主要功能手写数字识别、模型训练与评估、单张图片预测、HTTP 接口调用、批量文件夹预测硬件要求CPU 可训练推荐 NVIDIA GPU 加速显存占用与 batch size 和模型规模相关本项目小模型显存占用通常较低实际以本机运行为准启动方式Python 脚本训练 Uvicorn 启动 API 服务接口能力提供 /predict 接口返回预测类别和置信度批量任务支持批量图片文件夹预测结果输出为 CSV适合场景神经网络入门教学、快速原型验证、算法基线上线、接口服务开发演示这个项目不是某个工业级开源框架而是一套教学与工程结合的基线实现。它的价值在于代码量少、依赖清晰、逻辑容易跟踪适合作为理解神经网络代码执行流程的起点。读者跑通后可以把同样的编程模式迁移到 CIFAR-10、Fashion-MNIST 或者自定义图像分类任务上。2. 神经网络核心概念与代码关系2.1 前馈神经网络与反向传播我们常说的前馈神经网络信息从输入层进入经过隐藏层的加权计算最后从输出层产生预测结果。训练过程中损失函数衡量预测值和真实标签的差距反向传播算法把这个差距从输出层往回传递逐层计算每个权重对损失的梯度再用优化器去更新权重。对应到代码里模型的前向计算由forward方法实现outputs model(images)完成一次前向推断。反向计算则是由三行代码配合完成optimizer.zero_grad()清空上次留下的梯度。loss.backward()根据损失自动计算所有参数梯度。optimizer.step()用梯度更新权重。很多刚入门的人会把这三行代码当成固定套路但其实每一行都有明确作用。如果漏掉zero_grad梯度会累加最终参数更新方向就会错乱。理解这三行代码神经网络训练循环就理解了一半。2.2 卷积神经网络为什么适合图像任务卷积神经网络和全连接网络的核心区别是引入了卷积核在图像上滑动提取局部特征。图像本身有很强的空间局部性相邻像素之间的关联紧密而距离很远的像素之间关系弱。全连接层会把每个像素都连接到下一层的每个神经元上参数数量爆炸而且很容易忽略局部结构。卷积层通过权重共享用同一个卷积核扫描整张图参数数量大幅下降同时能提取边缘、纹理、形状等层级特征。本文的 SimpleCNN 由两个卷积块和一个全连接分类头组成。第一个卷积层把 1 通道灰度图升维成 32 通道特征图第二个卷积层再把特征图扩展到 64 通道每一层后面都接 ReLU 激活函数和最大池化。池化层的作用是降采样保留主要特征的同时减少计算量。最后把特征图展平通过两个全连接层输出 10 个类别的得分。2.3 循环神经网络与序列任务如果输入是时间序列、文本、语音这类带顺序关系的数据全连接网络和普通卷积网络就不好直接处理了。循环神经网络RNN每个时间步共享同一套权重把上一步的隐藏状态作为当前步输入的一部分从而让网络具备“记忆”能力。比如预测一段文本的下一个词、判断一段音频的情感、对时间序列做回归这些场景通常会用 RNN、LSTM 或 GRU 来建模。本文主要演示的是图像分类任务所以代码集中在卷积网络上。但从编程角度看RNN 的模型定义、训练循环、反向传播和保存加载流程与 CNN 是一致的。真正掌握本文的训练管线后再切换到 RNN 或者其他网络结构只是替换模型定义和数据加载部分而已。2.4 网络结构选择速览网络类型擅长任务典型输入关键代码对应前馈网络表格数据分类回归一维特征向量nn.Linear堆叠卷积神经网络图像分类检测分割二维图像nn.Conv2dnn.MaxPool2d循环神经网络文本序列时间序列变长序列数据nn.RNN、nn.LSTM图神经网络社交网络分子结构图结构数据依赖 PyTorch Geometric 等扩展库选网络结构的第一原则是匹配数据形态而不是一味追求复杂模型。图像任务先试 CNN文本任务先试 Transformer 或 RNN 系模型表格任务往往一个带 BatchNorm 和 Dropout 的多层感知机就已经够用。本次实操为了让代码容易在 CPU 上复现选择了 MNIST 小型 CNN 的组合。3. 适用场景与使用边界3.1 这套代码适合谁最直接的读者是刚学完神经网络基本原理、想看到实际代码的开发者。MNIST 数据集只有 28 × 28 像素单张图很小CPU 上也能够在几分钟内完成一个 epoch 的训练非常适合用来验证训练管线是否正常。项目代码量控制在可读完的范围每一段都能和训练流程对应起来调试起来也容易定位问题。工程侧的应用场景是快速搭建基线模型。比如一个图像分类需求刚提出来设计师还不确定最终数据规模和标注质量先用 MNIST 风格的代码结构跑通一套最小流程确认数据加载、模型训练、模型保存、接口调用四条链路都能工作再替换成真实数据集和更强的网络结构。这种从简单到复杂的推进方式比一上来就堆大模型更可控。3.2 不适合什么场景这套代码不适合直接部署为高并发生产服务。FastAPI 接口示例解决的是“能不能调通”的问题不是“抗压能力有多强”。如果需要支撑高并发、低延迟、GPU 资源调度、模型热更新、多模型版本管理需要引入更完整的推理服务框架。另外MNIST 模型只认识 28 × 28 的灰度手写数字。把它直接拿去做自然场景下的车牌识别、印刷体识别或其他复杂任务效果没有保障。严格来说一个模型只能处理它训练数据分布内的样本跨分布使用时必须重新评估。3.3 版权隐私与合规边界训练数据、测试数据、用户上传的图片都需要确认来源合法且具备使用授权。MNIST 数据集本身是公开研究数据集但不要默认所有公开数据都能随意商用。涉及人脸、证件、医疗影像等敏感图像时必须在数据脱敏、用户授权、隐私保护三个层面做好合规控制。接口服务如果部署在公网需要加访问鉴权不能裸奔暴露在公网上防止被恶意调用刷流量或者上传违规内容。识别结果如果要用到自动化决策场景要建立人工复核机制不能用一个小模型的置信度直接决定重要结果。4. 环境准备与依赖安装4.1 操作系统与 Python 版本项目代码在 Windows、Linux、macOS 上都可以运行。最顺手的是 Linux 服务器环境GPU 驱动和 CUDA 版本管理都更直观Windows 玩家用命令行工具运行三段脚本也没有障碍。Python 建议使用 3.9 及以上版本推荐 3.10 或 3.11过旧版本容易遇到依赖库不支持的问题。检查 Python 版本python --version如果系统里同时存在多个 Python 版本建议为这个项目单独创建虚拟环境避免依赖互相污染。4.2 创建虚拟环境Windows 命令行python -m venv venv venv\Scripts\activateLinux / macOSpython -m venv venv source venv/bin/activate激活后命令行前缀会出现(venv)说明当前处于虚拟环境中。后续所有 pip 安装和 Python 命令都在这个环境内执行。4.3 安装 PyTorchCPU 版本可以直接通过 pip 安装pip install torch torchvisionNVIDIA GPU 用户需要到 PyTorch 官网选择与本地 CUDA 版本匹配的安装命令。判断本地 CUDA 环境可以用nvidia-sminvidia-smi 显示的 CUDA Version 本身是驱动支持的版本不完全等于 PyTorch 运行时需要的 CUDA 版本但可以作为选版本的参考。安装完成后用 Python 验证 GPU 是否可用import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0) if torch.cuda.is_available() else CPU only)如果torch.cuda.is_available()返回 False通常有几种可能安装了 CPU 版 PyTorch、NVIDIA 驱动版本太旧、CUDA 运行时库不匹配。先检查 PyTorch 是否能看到当前驱动支持的 CUDA。4.4 安装其他依赖本文还需要 FastAPI、Uvicorn、Pillow、Requests写入需求文件torch torchvision fastapi uvicorn python-multipart pillow requests安装命令pip install -r requirements.txtpython-multipart是 FastAPI 接收文件上传时必须的依赖漏装会导致上传接口报错。4.5 数据集准备首次运行训练脚本时Torchvision 会尝试从网络自动下载 MNIST 数据集。如果下载速度很慢或者超时可以手动从可访问的镜像站点下载四个.gz文件train-images-idx3-ubyte.gztrain-labels-idx1-ubyte.gzt10k-images-idx3-ubyte.gzt10k-labels-idx1-ubyte.gz把这些文件放到项目的data/MNIST/raw/目录下然后重新运行训练脚本。Torchvision 检测到目录中已有原始文件就不会重复下载。数据文件归档名最好保持原样Torchvision 内部用固定文件名去读取。5. 神经网络代码结构设计与实现5.1 项目目录结构neural-network-demo/ ├── requirements.txt ├── train.py ├── infer.py ├── api_server.py ├── batch_predict.py ├── data/ │ └── MNIST/ └── outputs/ └── mnist_cnn.pthtrain.py数据加载、模型定义、训练循环、模型保存。infer.py加载模型对单张图片做预测。api_server.py启动 HTTP 接口服务通过前端或脚本远程调用。batch_predict.py遍历文件夹内所有图片批量预测并输出 CSV。outputs/存放训练好的权重文件。教学场景下每个脚本独立可运行微信传给别人的时候不用考虑复杂 import 路径问题。正式项目里建议把模型定义抽到model.py把工具函数抽到utils.py避免重复代码。这里为了保证读者直接复制就能跑做了一个简化设计。5.2 数据加载与预处理train.py的数据加载部分import time import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms def load_data(batch_size128): transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_dataset datasets.MNIST( root./data, trainTrue, downloadTrue, transformtransform ) test_dataset datasets.MNIST( root./data, trainFalse, downloadTrue, transformtransform ) train_loader DataLoader( train_dataset, batch_sizebatch_size, shuffleTrue, num_workers2 ) test_loader DataLoader( test_dataset, batch_size256, shuffleFalse, num_workers2 ) return train_loader, test_loader关键点是transforms.ToTensor()把 PIL 图片或者 ndarray 转成 PyTorch Tensor同时把像素值缩放到 0 到 1 之间transforms.Normalize再用 MNIST 数据集的均值和标准差做标准化让输入分布更稳定。shuffleTrue保证训练时每个 batch 的数据顺序都是打乱的避免模型学到样本顺序上的假关联。测试集不需要打乱只需要顺序按批次遍历。5.3 CNN 模型定义train.py里的模型结构class SimpleCNN(nn.Module): def __init__(self): super(SimpleCNN, self).__init__() self.conv1 nn.Conv2d(1, 32, kernel_size3) self.conv2 nn.Conv2d(32, 64, kernel_size3) self.pool nn.MaxPool2d(2) self.flatten nn.Flatten() self.fc1 nn.Linear(64 * 5 * 5, 128) self.relu nn.ReLU() self.dropout nn.Dropout(0.2) self.fc2 nn.Linear(128, 10) def forward(self, x): x self.pool(self.relu(self.conv1(x))) x self.pool(self.relu(self.conv2(x))) x self.flatten(x) x self.relu(self.fc1(x)) x self.dropout(x) x self.fc2(x) return x输入是 1 × 28 × 28 的灰度图。第一层卷积使用 32 个 3 × 3 卷积核输出特征图尺寸变成 32 × 26 × 26池化后变成 32 × 13 × 13。第二层卷积输出 64 × 11 × 11池化后变成 64 × 5 × 5。展平后得到 64 × 5 × 5 1600 维向量经过 128 维全连接层和 Dropout最后映射到 10 类。nn.Flatten的作用是把多维特征图拉成一维向量交给全连接层处理。Dropout 在训练时随机把一部分神经元输出置零起到正则化作用测试时自动关闭不需要手动切换。5.4 训练与评估循环训练一个 epochdef train_one_epoch(model, loader, criterion, optimizer, device): model.train() total_loss 0.0 correct 0 total 0 for images, labels in loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() * images.size(0) preds outputs.argmax(dim1) correct (preds labels).sum().item() total labels.size(0) return total_loss / total, correct / total评估函数def evaluate(model, loader, criterion, device): model.eval() total_loss 0.0 correct 0 total 0 with torch.no_grad(): for images, labels in loader: images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) total_loss loss.item() * images.size(0) preds outputs.argmax(dim1) correct (preds labels).sum().item() total labels.size(0) return total_loss / total, correct / total评估阶段必须用torch.no_grad()包裹推理过程告诉 PyTorch 不需要计算梯度。这样既能省显存也能加快计算速度。model.eval()和model.train()的切换要养成习惯因为 Dropout 和 BatchNorm 在两种模式下的行为不同。本项目代码如果用训练模式去做推理Dropout 仍然生效预测结果会带有随机性导致同样的图片每次预测结果不一样这是新手容易踩的坑。主训练函数def main(): device torch.device(cuda if torch.cuda.is_available() else cpu) print(Using device:, device) train_loader, test_loader load_data(batch_size128) model SimpleCNN().to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr0.001) epochs 5 for epoch in range(1, epochs 1): start time.time() train_loss, train_acc train_one_epoch( model, train_loader, criterion, optimizer, device ) test_loss, test_acc evaluate(model, test_loader, criterion, device) cost time.time() - start print( fEpoch {epoch}/{epochs} | ftrain_loss{train_loss:.4f} train_acc{train_acc:.4f} | ftest_loss{test_loss:.4f} test_acc{test_acc:.4f} | ftime{cost:.1f}s ) torch.save(model.state_dict(), outputs/mnist_cnn.pth) print(Model saved to outputs/mnist_cnn.pth) if __name__ __main__: main()CrossEntropyLoss在 PyTorch 里已经包含 Softmax 计算所以模型的最后一层不需要额外加 Softmax 激活。训练阶段最好直接输出原始 logits传给损失函数计算需要概率时推理阶段再调用torch.softmax手动转换。这样数值稳定性最好。训练过程观察点train_acc是否逐 epoch 上升。test_acc是否同步上升。train_loss是否持续下降。如果训练准确率很高但测试准确率停滞说明过拟合。如果训练准确率一直上不去优先检查数据增强、学习率、模型容量。从常见实践来看MNIST 这种简单数据集上用上面这个模型训练 5 个 epoch 通常能达到 98% 以上的测试准确率。具体数值会因为随机种子、数据加载顺序、优化器状态不同而有小幅度波动以本机运行为准。5.5 单张图片推理训练完成后用infer.py把模型权重加载回来对任意手写数字图片做预测。这种“训练一次、反复推理”的方式也是生产环境中使用模型的常规方式。import torch from PIL import Image from torchvision import transforms class SimpleCNN(torch.nn.Module): def __init__(self): super(SimpleCNN, self).__init__() self.conv1 torch.nn.Conv2d(1, 32, kernel_size3) self.conv2 torch.nn.Conv2d(32, 64, kernel_size3) self.pool torch.nn.MaxPool2d(2) self.flatten torch.nn.Flatten() self.fc1 torch.nn.Linear(64 * 5 * 5, 128) self.relu torch.nn.ReLU() self.dropout torch.nn.Dropout(0.2) self.fc2 torch.nn.Linear(128, 10) def forward(self, x): x self.pool(self.relu(self.conv1(x))) x self.pool(self.relu(self.conv2(x))) x self.flatten(x) x self.relu(self.fc1(x)) x self.dropout(x) x self.fc2(x) return x def preprocess_image(image_path): transform transforms.Compose([ transforms.Grayscale(num_output_channels1), transforms.Resize((28, 28)), transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) image Image.open(image_path).convert(L) return transform(image).unsqueeze(0) def predict_image(model, image_path, device): tensor preprocess_image(image_path).to(device) model.eval() with torch.no_grad(): logits model(tensor) prob torch.softmax(logits, dim1) pred torch.argmax(prob, dim1).item() confidence prob[0][pred].item() return pred, confidence if __name__ __main__: device torch.device(cuda if torch.cuda.is_available() else cpu) model SimpleCNN().to(device) model.load_state_dict(torch.load(outputs/mnist_cnn.pth, map_locationdevice)) image_path test_digit.png pred, conf predict_image(model, image_path, device) print(fPredicted: {pred}, Confidence: {conf:.4f})map_locationdevice是一个很容易被忽略的参数。如果模型在 GPU 上训练保存的权重默认在 GPU 显存里换一台只有 CPU 的机器加载时必须指定map_locationcpu否则会报设备不匹配的错误。6. 训练效果与验证方法6.1 日志指标解读运行python train.py后输出日志如下Using device: cpu Epoch 1/5 | train_loss0.2103 train_acc0.9385 | test_loss0.0764 test_acc0.9761 | time23.4s Epoch 2/5 | train_loss0.0741 train_acc0.9771 | test_loss0.0562 test_acc0.9827 | time22.9s Epoch 3/5 | train_loss0.0521 train_acc0.9837 | test_loss0.0451 test_acc0.9851 | time23.1s Epoch 4/5 | train_loss0.0415 train_acc0.9872 | test_loss0.0398 test_acc0.9874 | time23.2s Epoch 5/5 | train_loss0.0340 train_acc0.9896 | test_loss0.0362 test_acc0.9882 | time23.0s时间消耗和具体硬件关系很大上面数据只是给一个数量级参考。CPU 上单 epoch 通常几十秒量级GPU 上会快很多尤其是 batch size 和图像尺寸都很小的情况下。训练集准确率略高于测试集准确率属于正常现象因为模型天然在见过的数据上表现更好如果这个差距拉得特别大说明过拟合严重。6.2 混淆矩阵进一步观察错误准确率只能反映整体正确比例看不出模型把哪两类数字搞混了。对多分类模型用混淆矩阵统计每一类的真实标签和预测标签分布能够快速定位问题。比如数字 4 经常被识别成 9、数字 7 经常被识别成 1这些错误类型在混淆矩阵里一眼就能看到。简单实现一个混淆矩阵可视化import numpy as np import matplotlib.pyplot as plt from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay # all_preds 和 all_labels 是评估时收集的预测与真实标签列表 cm confusion_matrix(all_labels, all_preds) disp ConfusionMatrixDisplay(confusion_matrixcm) disp.plot(cmapBlues) plt.title(Confusion Matrix on MNIST Test Set) plt.savefig(confusion_matrix.png, dpi150)如果没有安装 scikit-learn可以手动创建一个 10 × 10 的零矩阵然后逐个样本累加cm[true_label][pred_label] 1。逻辑不复杂但 scikit-learn 的接口更标准适合做基线分析。生成混淆矩阵图片后重点是看对角线元素占比以及哪些非对角元素显著偏高。6.3 过拟合与欠拟合判断欠拟合表现为训练准确率和测试准确率都低通常需要对模型加容量、增加训练轮数、调低学习率或者给数据做更多增强。过拟合表现为训练准确率很高但测试准确率下降应对策略包括增加训练数据、加大 Dropout、加权重衰减、提前停止训练。MNIST 数据量足够大、任务足够简单上面这个小模型基本不会出现过拟合所以 Dropout 在这里更多是演示标准训练范式的作用。如果换成自己的业务数据训练集只有几千张模型又是一个大网络过拟合几乎是必然发生的。建议在训练脚本里同时打印 train loss 和 test loss不要只盯准确率。Loss 的差距比准确率差距更能反映模型泛化状态。7. 接口 API 与批量任务7.1 启动 FastAPI 推理服务模型训练完后要把它变成可供外部程序调用的服务。api_server.py实现一个图片上传接口接收请求后返回预测类别和置信度。import io import torch from fastapi import FastAPI, UploadFile, File from PIL import Image from torchvision import transforms class SimpleCNN(torch.nn.Module): # 模型定义与训练脚本保持一致 pass app FastAPI(titleMNIST Inference API) model SimpleCNN() model.load_state_dict(torch.load(outputs/mnist_cnn.pth, map_locationcpu)) model.eval() transform transforms.Compose([ transforms.Grayscale(num_output_channels1), transforms.Resize((28, 28)), transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) app.post(/predict) async def predict(file: UploadFile File(...)): image_bytes await file.read() image Image.open(io.BytesIO(image_bytes)).convert(L) tensor transform(image).unsqueeze(0) with torch.no_grad(): logits model(tensor) prob torch.softmax(logits, dim1) pred int(torch.argmax(prob, dim1).item()) confidence float(prob[0][pred].item()) return {label: pred, confidence: confidence}以上代码中pass处应替换为模型定义为节约篇幅用了占位。实际运行时必须保持一致不然加载权重会报结构不匹配。启动服务uvicorn api_server:app --host 127.0.0.1 --port 8000启动后终端会显示访问地址本地浏览器打开http://127.0.0.1:8000/docs可以查看 Swagger 文档直接测试接口。7.2 用 Python 调用接口接口服务跑起来后用下面的脚本测试import requests url http://127.0.0.1:8000/predict files {file: (digit.png, open(digit.png, rb), image/png)} response requests.post(url, filesfiles, timeout30) print(response.json())预期输出{label: 7, confidence: 0.9992}label 表示预测的数字confidence 表示模型对这个结果的置信度。接口返回这种结构化 JSON方便后续接到其他业务系统里。7.3 批量图片推理脚本实际使用场景中经常需要对一个文件夹里大量图片做推理。逐张调用 HTTP 接口当然可行但本机批量处理时直接加载模型跑更高效。batch_predict.py实现文件夹遍历、逐图预测、CSV 导出。import csv import glob import os import torch from PIL import Image from torchvision import transforms class SimpleCNN(torch.nn.Module): pass def preprocess_image(image_path): transform transforms.Compose([ transforms.Grayscale(num_output_channels1), transforms.Resize((28, 28)), transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) image Image.open(image_path).convert(L) return transform(image).unsqueeze(0) def batch_predict(model, image_dir, output_csv, device): model.eval() results [] image_paths sorted(glob.glob(os.path.join(image_dir, *.png))) for path in image_paths: tensor preprocess_image(path).to(device) with torch.no_grad(): logits model(tensor) prob torch.softmax(logits, dim1) pred int(torch.argmax(prob, dim1).item()) confidence float(prob[0][pred].item()) results.append({ image: os.path.basename(path), label: pred, confidence: confidence }) print(f{os.path.basename(path)} - {pred} ({confidence:.4f})) with open(output_csv, w, newline, encodingutf-8) as f: writer csv.DictWriter(f, fieldnames[image, label, confidence]) writer.writeheader() writer.writerows(results) print(fDone. {len(results)} images, output - {output_csv}) if __name__ __main__: device torch.device(cuda if torch.cuda.is_available() else cpu) model SimpleCNN().to(device) model.load_state_dict(torch.load(outputs/mnist_cnn.pth, map_locationdevice)) batch_predict(model, ./test_images, ./outputs/predictions.csv, device)批量脚本需要加日志和失败重试机制尤其当图片文件很多、某些图格式异常时。上面这个版本是“遇到异常直接退出”的简化版工程上建议每个文件单独用 try except 包住单张失败记录到error.log不影响其他图片继续处理。8. 资源占用与性能观察8.1 显存占用怎么看训练时观察 GPU 显存最直接的方式是nvidia-smi -l 1-l 1表示每秒刷新一次。在训练脚本运行的同时打开另一个终端窗口执行这个命令可以看到 Python 进程占用的显存大小和 GPU 利用率。在 PyTorch 代码内部记录显存峰值print(torch.cuda.max_memory_allocated() / 1024**2, MB)这个数值表示当前程序在 GPU 上分配的峰值显存。它对调优 batch size 很实用如果显存接近上限就先减小 batch size再考虑换更小的模型。8.2 影响资源占用的因素batch size 越大单次前向和反向计算需要的显存越高。图片分辨率越大特征图占用的显存越大。网络层数越多、通道数越多参数和中间激活值越多。训练阶段显存占用通常高于推理阶段因为反向传播需要保存中间梯度。使用混合精度训练可以减少显存占用并提升速度但需要显卡支持。MNIST 是 28 × 28 的小图这个项目如果用 GPU 训练显存占用会很克制。实际跑自己的高清业务数据时显存压力会成倍上升尤其要关注中间特征图的尺寸而不是只盯模型参数量。8.3 CPU 与 GPU 训练差异CPU 可以完成同样的训练任务PyTorch 在 CPU 上使用多线程加速但和 GPU 比差距悬殊。MNIST 这种小任务上GPU 的优势更多体现在大数据集、大模型、大 batch size 场景。CPU 训练时如果发现电脑卡顿可以先减小num_workers降低数据加载线程对系统资源的抢占GPU 训练时则要关注数据加载速度避免 GPU 空转等待数据。用torch.cuda.is_available()判断设备后把模型和数据显式.to(device)代码就可以自动适配 CPU 和 GPU。这也是 PyTorch 开发的标准写法。上面所有脚本都对设备做了统一处理换机器不需要改逻辑。8.4 降低显存占用的通用方法调小 batch size。减少输入图片尺寸。减少网络通道数。使用torch.no_grad()包裹推理代码避免保存计算图。训练时定期清理不再使用的计算图缓存。使用梯度累积用小 batch size 模拟大 batch size 的效果。模型太大时考虑 DDP 分卡或模型并行但小项目不需要上这么重。9. 常见问题与排查方法问题现象可能原因排查方式解决方案MNIST 下载非常慢或超时默认源访问不稳定检查网络与data/MNIST/raw/目录手动下载四个 gz 文件放入 raw 目录后重试torch.cuda.is_available()返回 False安装的是 CPU 版 PyTorch或驱动不匹配pip list查看 torch 版本运行nvidia-smi到 PyTorch 官网按 CUDA 版本重装训练时显存不足 OOMbatch size 过大或显卡显存太小查看日志中的 CUDA out of memory 报错调小 batch size降分辨率用梯度累积模型加载权重时报 size mismatchload_state_dict时的模型结构和训练时不一致检查模型定义代码是否有改动确认模型类定义与训练时完全一致接口调用一直 500文件上传格式不对或依赖缺失查看 Uvicorn 终端堆栈信息确认已安装python-multipart检查图片格式推理结果全部是同一个类别数据预处理不一致对比训练和推理的 transform统一用 Grayscale、Resize、Normalize 参数端口被占用上一次服务未退出或端口冲突检查终端报错Windows 用netstat -ano换一个端口或者杀掉占用进程后重启预测置信度全部很低图片和训练数据分布差距大检查输入图片是否清晰、尺寸是否正确对图片做预处理增强或收集更接近的数据重新训练每轮训练准确率波动大学习率过高或数据加载随机性观察 loss 曲线调低学习率固定随机种子固定随机种子的方法import random import numpy as np import torch def set_seed(seed42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed)在训练前调用一次结果的可复现性会明显提升。数据加载过程的随机增强仍然可能带来微小波动但对 MNIST 这种简单数据集影响不大。如果.pth权重文件丢失或者想重新训练直接在项目根目录运行python train.py即可所有中间产物都会重建。outputs/predictions.csv也可以随时删除重跑不影响模型文件本身。10. 最佳实践与总结10.1 工程化建议第一数据、代码、输出分目录管理。本项目虽然简单但目录边界清晰data/放原始数据项目根部放脚本outputs/放模型和结果。真实项目建议再加一个configs/目录存放超参数配置、一个logs/目录存放训练日志。第二第一次跑训练先不追求精度用小 batch size 跑几个 batch 验证代码流是否通畅。比如在训练循环里加一个if batch_idx 20: break的临时逻辑确认前向传播、反向传播、权重更新均无报错之后再放开完整训练。第三接口服务要加访问限制。--host 127.0.0.1表示只监听本机请求适合本地调试。如果需要在局域网内访问改成0.0.0.0后一定要配防火墙策略或 Token 鉴权。生产环境建议在反向代理层做请求体大小限制、速率限制和身份验证。第四批量任务必须幂等。重复执行batch_predict.py不应产生脏数据建议每次运行前清空目标 CSV或者在文件名里加时间戳。对超时或失败的图片单独记录失败原因不要把整个任务回滚掉。10.2 优先验证的三个点先验证设备识别跑python -c import torch; print(torch.cuda.is_available())确认训练脚本最终落在 CPU 还是 GPU 上。再验证训练闭环完整跑一个 epoch关注 loss 是否下降、准确率是否高于随机猜测。MNIST 二分类基线准确率约 50%十分类随机猜约 10%如果第一次 epoch 训练准确率就到了 90% 以上说明数据和模型链路正常。最后验证接口链路启动 FastAPI 后用一张真实图片调用/predict确认返回的 label 和 confidence 字段格式正确。接口能跑通这套代码就可以接到自己的工具链里继续扩展。10.3 可以继续扩展的方向这套代码的骨架可以直接迁移到其他图像分类数据集。把SimpleCNN的输入通道改成 3把最后的全连接层输出改成目标类别数数据加载改成自己的图片目录结构训练脚本其余部分基本不动。训练部分可以增加学习率调整策略、模型保存时同时保存 optimizer 状态、加入早停机制。推理部分可以尝试用 ONNX 导出模型用 ONNX Runtime 加速 CPU 推理或者用 TensorRT 在 GPU 上做更极致的性能优化。接口部分可以增加模型版本号和请求 ID方便排障和灰度发布。从教学角度看下一次升级可以把 CNN 换成 LSTM 试试对序列数据的处理也可以把模型换成预训练 ResNet 做迁移学习对比不同网络结构的精度和资源差异。理解这个最小项目的完整链路之后再往上叠功能每一步都有明确方向。