简介:一套面向计算机专业课程设计、期末大作业与毕业设计场景的CNN手写数字识别项目,基于Python搭建卷积神经网络训练与识别流程,代码经导师指导并获评高分,初学者可下载后直接运行复现。资源共24个文件,包含5个Python源码脚本(模型训练、登录主界面、图片识别等)、6份docx报告文档(需求分析、系统设计、测试用例与需求验证)、MNIST相关数据集及README说明,压缩包约31.45MB,按代码、报告、数据分区,方便对照查阅。目前已有89人学习,资料从需求规格到测试验收完整闭环,既能支持快速交付可演示的大作业,也能帮助理解CNN原理、图像预处理与模型评估的实践链路。
1. CNN 手写数字识别:一份能直接跑的课程设计,到底给了你什么?
期末周最磨人的往往不是卷积公式,而是答辩前一晚发现训练好的模型在别人电脑上起不来。这份基于 Python 实现的 CNN 卷积神经网络手写数字识别项目,好就好在它是按课程设计标准打包的:训练模型.py、数字图片识别.py、登录主界面.py、完整代码.py 分得清清楚楚,外加 mnist_pic.zip 数据集和整套需求分析、系统设计、测试用例、实验报告文档。换句话说,你拿到的不是一段孤零零的神经网络代码,而是一个从需求分析到测试验收的闭环。适合两类人:一是计算机专业准备期末大作业、课程设计想省时间的学生,二是想完整跑通 CNN 训练与推理链路的学习者。项目评审分 99 分,代码能跑,文档能交,照着下面的步骤复现,至少少走三天弯路。
2. 把项目拆开看:目录结构、代码职责与运行顺序
先别急着 pip install。拿到压缩包第一步,是搞清楚哪些文件是主角,哪些是配角。这份项目的目录不算大,但文件不少,而且命名一眼就能看出是给课程设计准备的:训练模型.py 负责训练,数字图片识别.py 负责推理,登录主界面.py 是花活,完整代码.py 是把所有东西串起来跑一遍,主要功能.py 则是公共函数封装。文档那一摞 docx 用于交作业,不影响运行。
2.1 压缩包里的文件各自负责什么
我习惯把它整理成一张职责表,免得运行时报错时不知道看谁。这份资源大致对应下表:
| 文件 | 作用 | 运行时机 |
|---|---|---|
| requirements.txt | 记录依赖库 | 第一步安装 |
| 训练模型.py | 定义CNN结构、训练、保存模型文件 | 第一次运行 |
| 数字图片识别.py | 加载模型、预测单张图片 | 训练之后 |
| 完整代码.py | 把训练+识别+展示串成一条流 | 汇报演示 |
| 主要功能.py | 封装预处理、预测等公共函数 | 被其他文件 import |
| 登录主界面.py | GUI 登录界面,课程设计展示用 | 演示时单独跑 |
| mnist_pic.zip | 手写数字图片数据集 | 训练或测试时读入 |
| README.md / README.en.md | 项目说明与运行入口 | 开始前必读 |
| 项目需求分析.docx 等 | 需求、设计、测试文档 | 交作业 |
从这份表可以反推出标准运行顺序:先装环境,再跑训练模型.py,模型文件生成后跑数字图片识别.py 验证,最后打开完整代码.py 看整体效果。登录主界面.py 是独立的演示壳,依赖前面的模型,但逻辑上跟 CNN 训练没有关系,别把它当成主入口。
2.2 先读 README 和依赖,再决定从哪个文件入手
很多同学解压后直接双击完整代码.py,结果报错就蒙了。我一般会先打开 README.md 和 requirements.txt。README 会写清楚运行环境和步骤,而 requirements.txt 决定了你能不能装出兼容的环境。这份项目里大概率是下面这类依赖声明:
tensorflow==2.10.0 numpy opencv-python Pillow matplotlib注意,我这里写的版本号是常见组合,实际以你压缩包里的 requirements.txt 为准。逻辑很简单:TensorFlow 负责 CNN 模型的训练和推理,numpy 处理矩阵运算,OpenCV 和 Pillow 负责图像读写与预处理,matplotlib 用来画训练曲线和混淆矩阵。参数上最需要留心的是 tensorflow 那行,因为它跟 Python 版本强绑定,后面我会单独说。
看完这两份文件,再决定从哪个文件入手就是顺理成章的事:如果只是作业演示,跑完整代码.py 就行;如果想看模型细节,读训练模型.py;如果想做单张图片测试,直接跑数字图片识别.py。不要一开始就钻进文档里看需求分析,那不是工程入口。
2.3 环境配置:Python版本、TensorFlow/Keras与requirements.txt的坑
项目能跑的前提是环境对得上。以 TensorFlow 2.10 这代版本为例,Python 3.8 或 3.9 最稳,3.10 偶尔会遇到依赖编译问题,3.12 则是经常直接没有对应的预编译轮子。我习惯在项目根目录建一个虚拟环境,不让全局环境被污染:
python -m venv venv venv\Scripts\activate pip install -r requirements.txt第一行创建虚拟环境,第二行激活,第三行安装依赖。如果你是 Linux 或 macOS,激活命令换成source venv/bin/activate。这里最值得强调的参数是 Python 版本,如果你机器上只有 Python 3.12,建议先安装 3.9 再继续,不要硬刚 pip 编译错误。
装完依赖后,运行下面这行确认 TensorFlow 能正常导入:
python -c "import tensorflow as tf; print(tf.__version__)"如果输出版本号,说明环境这关过了。如果卡在某个 DLL 或者 AVX 错误,多半是版本不匹配,卸载重装对应版本,不要满世界找补丁。还有一个经常被忽略的组合问题是 opencv-python 和 numpy 的版本匹配,装完 import cv2 报错说找不到npy_intp类型时,我一般会先把 numpy 降到 1.23.x 再重装 opencv,这算是我踩过最多次的血泪经验。
环境搞定后,看一眼数据集的放置位置:训练脚本里如果写的是mnist.load_data(),它会自动下载;如果写的是mnist_pic.zip,那就要先解压并保证相对路径正确。运行训练时,如果不想看满屏 INFO 日志,可以先设一个环境变量:
export TF_CPP_MIN_LOG_LEVEL=2 python 训练模型.pyTF_CPP_MIN_LOG_LEVEL=2的意思是只显示 ERROR 级别日志,把 INFO 和 WARNING 都关掉。这样做的好处是训练过程清爽,报错信息一眼就能看到。这里面最容易忽略的坑是:不要从资源管理器双击运行 py 文件,尽量在项目目录下的终端里执行,否则相对路径会基于当前工作目录解析,数据集目录一变就找不到。
再补一句:这份资源里的代码文件命名虽然直白,但运行顺序不同,报错信息也完全不同。先把训练模型.py 跑通,再管识别,顺序别反。很多同学一上来就点登录主界面.py,结果后端模型还没生成,界面闪退,这属于预期内现象,不是代码坏了。
3. 核心网络与训练流程:MNIST 数据怎么进 CNN
手写数字识别最经典的基准就是 MNIST:60000 张 28×28 灰度训练图,10000 张测试图,每张图上一个 0 到 9 的数字。这份项目里 mnist_pic.zip 是打包好的图片集,但很多写法是直接调用 Keras 内置的 mnist 接口。如果代码里用的是from tensorflow.keras.datasets import mnist,那么第一次运行会自动从网络下载到用户目录,之后就一直复用。搞清楚数据和模型是怎么接上的,比背一句“CNN 能识别手写数字”有用得多。
3.1 数据加载与预处理
预处理的关键是把普通图片数组变成 CNN 期望的四维张量。CNN 的输入形状是(样本数, 高, 宽, 通道数),而 mnist.load_data() 返回的 x_train 是三维(60000, 28, 28),没带通道这一维。常见处理代码长这样:
from tensorflow.keras.datasets import mnist from tensorflow.keras.utils import to_categorical (x_train, y_train), (x_test, y_test) = mnist.load_data() x_train = x_train.reshape(-1, 28, 28, 1).astype("float32") / 255.0 x_test = x_test.reshape(-1, 28, 28, 1).astype("float32") / 255.0 y_train = to_categorical(y_train, 10) y_test = to_categorical(y_test, 10)第一处reshape(-1, 28, 28, 1)里,-1 是让 numpy 自动算成 60000,后面的 1 表示灰度单通道。astype("float32")是为了避免内存占用翻倍,也防止某些环境用 float64 训练变慢。除以 255 是最常见的归一化,把像素范围压到 0 到 1,这对激活函数和梯度传播都友好。如果训练脚本里没有这一步,loss 很容易飘。
标签那边,to_categorical(y_train, 10)把标量 3 变成[0, 0, 1, 0, 0, 0, 0, 0, 0, 0],这是为了配合分类交叉熵损失函数。如果不做 one-hot,模型会把数字当回归问题来学,准确率会掉一截。
如果这份项目里用的是本地 mnist_pic.zip 而不是内置 mnist,也不要慌。把 zip 解压后按 0 到 9 分文件夹存放,用 OpenCV 循环读取即可:
import cv2 import os import numpy as np def load_from_folder(data_dir="mnist_pic"): X, y = [], [] for label in range(10): folder = os.path.join(data_dir, str(label)) for fname in os.listdir(folder): img = cv2.imread(os.path.join(folder, fname), cv2.IMREAD_GRAYSCALE) img = cv2.resize(img, (28, 28)) X.append(img) y.append(label) X = np.array(X).reshape(-1, 28, 28, 1).astype("float32") / 255.0 return X, np.array(y)这段代码的关键是cv2.IMREAD_GRAYSCALE强制读成单通道,避免彩色三通道进来。resize保证每张图都是 28×28。如果你发现本地图片集训练效果差,先检查是不是有彩色图混进来,再检查文件夹命名和 label 是否对应。
3.2 模型结构:卷积层、池化层、全连接层参数设计
经典 LeNet 风格的网络在这份作业里完全够用。它的主干是两个卷积-池化块,再接全连接层。卷积层自动提取局部特征,池化层压缩尺寸,全连接层组合成全局特征,最后的 softmax 输出十个类概率。具体可以写成一个 Sequential:
from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Conv2D, MaxPooling2D, Flatten, Dense, Dropout model = Sequential([ Conv2D(32, (3, 3), activation="relu", input_shape=(28, 28, 1)), MaxPooling2D((2, 2)), Conv2D(64, (3, 3), activation="relu"), MaxPooling2D((2, 2)), Flatten(), Dense(128, activation="relu"), Dropout(0.25), Dense(10, activation="softmax") ]) model.summary()参数上,第一个 Conv2D 用 32 个 3×3 卷积核,input_shape必须和预处理后的数据形状一致,如果前面写(28, 28, 1)这里写成(784,),第一层就报错。MaxPooling2D 把 26×26 压成 13×13,参数量立刻少了四分之三。第二个卷积核翻倍到 64,是因为越往高层特征越抽象,通道数适当地增加能提升表达能力,但翻到 128 在没有数据增强时容易过拟合。全连接层 128 是个平衡点,太大训练慢,太小准确率上不去。Dropout 放全连接后面,随机关闭 25% 的神经元,是课程设计里最不亏的防过拟合手段。
model.summary()是免费的检查工具。跑一下就能看到每层输出的 shape 和参数量,如果后面层数计算对不上,多半是 Flatten 前的尺寸和你预期不一致。我经常在答辩时让同学现场跑这句,能看到输出 1600 维,比干背结构可信。
3.3 训练模型.py 里那些值得抄的训练参数
训练这步直接关系到报告里最醒目的那个准确率。编译和 fit 的写法是这份项目里最值得抄的部分:
model.compile( optimizer="adam", loss="categorical_crossentropy", metrics=["accuracy"] ) history = model.fit( x_train, y_train, batch_size=128, epochs=10, validation_split=0.1, verbose=1 )优化器选 adam 是因为它对学习率不敏感,新手不用手动调学习率,梯度下降的收敛也稳。loss 用categorical_crossentropy,它要求标签是 one-hot;如果标签是整数,就要换sparse_categorical_crossentropy。这两个混淆是课程设计里的高频报错点,原理想清楚就能一眼看出问题。metrics 里写 accuracy 是为了在训练过程中直接看准确率,画图也有数据来源。
batch_size=128对于 60000 张图来说,一个 epoch 约 470 个 batch,显存占用适中,速度也快。如果你改成 16,模型更新次数暴增,训练时间拉长,但准确率不一定更高。epochs=10在 MNIST 上通常能到 99% 以上,再多到了后期 val_accuracy 基本不再涨,纯属烧时间。validation_split=0.1表示在训练集内部再留 6000 张当验证集,省得手动切数据。它是 fit 里很实用的参数,会给报告提供一条完整的 val_acc 曲线。
训练结束后不要忘了保存模型:
model.save("mnist_cnn.h5")注意,Keras 3 的默认格式已经不是 h5,如果你用的 TensorFlow 版本较新,可能建议直接保存为mnist_cnn.keras。但很多旧代码依然写 h5,能加载就行。保存路径建议直接用项目根目录相对路径,这样数字图片识别.py 加载时不用改来改去。
3.4 训练过程中的可视化与保存
报告文档里那几张准确率曲线图,其实就是从 history 里画的。fit 的返回值记录了每个 epoch 的 loss 和 accuracy,直接拿 matplotlib 画线:
import matplotlib.pyplot as plt plt.plot(history.history["accuracy"], label="train_acc") plt.plot(history.history["val_accuracy"], label="val_acc") plt.xlabel("epoch") plt.ylabel("accuracy") plt.legend() plt.savefig("acc_curve.png", dpi=150)history.history是一个字典,键就是 compile 里 metrics 名加 loss。dpi=150导出的图片放在报告里足够清晰。很多同学到答辩时才发现曲线图太糊,其实只要这一句就能解决。
除了准确率,建议顺手保存一条 loss 曲线。答辩老师常问的是“你有没有过拟合”,你有两张曲线就能从走势回答:训练 loss 持续下降而验证 loss 回升,就是过拟合信号。这是课程设计报告里最容易拿分的细节,也最能体现你不是只跑了别人代码。
4. 把模型跑起来:数字图片识别.py 与登录主界面的完整链路
训练完拿到模型文件,下一步是把它接到识别和展示层。很多课程设计卡在这一步:训练时好好的,换到推理脚本就一堆报错。原因是训练和推理对数据的处理方式必须完全一致,一个像素的差异都会影响结果。这一章就把登录主界面、图片识别和完整代码之间的调用关系拆开。
4.1 登录主界面.py:先说清楚它和识别不是一回事
登录主界面.py 在课程设计里承担的是“系统外壳”。代码本身是 Tkinter 或 PyQt 的窗体,用户名密码验证通过后打开主窗口。不要期待它参与任何卷积计算。常见写法如下:
import tkinter as tk from tkinter import messagebox def check_login(): if user_var.get() == "admin" and pwd_var.get() == "123456": win.destroy() import complete_code else: messagebox.showerror("登录失败", "用户名或密码错误") win = tk.Tk() user_var = tk.StringVar() pwd_var = tk.StringVar() tk.Label(win, text="用户名").pack() tk.Entry(win, textvariable=user_var).pack() tk.Label(win, text="密码").pack() tk.Entry(win, textvariable=pwd_var, show="*").pack() tk.Button(win, text="登录", command=check_login).pack() win.mainloop()代码逻辑很直接:两个 StringVar 接收输入,按钮触发check_login()。注意show="*"是密码输入框的掩码显示。这里要明确一点:密码写死在代码里是为了演示,真实系统不可能这么干。所以你在报告里不要吹自己做了安全认证,就说“模拟登录流程”即可。
import complete_code是常见的跳转写法,登录成功后销毁登录窗口,进入主程序。如果 complete_code.py 里有训练代码,登录后会自动开始训练或加载模型,所以这个文件其实是整个系统的大门。我见过有人把完整代码.py 当成纯前端文件去读,结果怎么都找不到界面逻辑,其实就是没理解它是个混合入口。
4.2 数字图片识别.py:从加载权重到输出预测
识别脚本更接近实际工程。加载模型后进行图像处理,然后预测。核心代码一般是这样的:
from tensorflow.keras.models import load_model import cv2 import numpy as np model = load_model("mnist_cnn.h5") def predict_digit(image_path, invert=False): img = cv2.imread(image_path, cv2.IMREAD_GRAYSCALE) img = cv2.resize(img, (28, 28)) if invert: img = 255 - img img = img.reshape(1, 28, 28, 1).astype("float32") / 255.0 probs = model.predict(img, verbose=0)[0] pred = int(np.argmax(probs)) conf = float(probs[pred]) return pred, confcv2.imread的第二个参数IMREAD_GRAYSCALE强制读成灰度图,这一步决定了后面 shape 一定是(28, 28)而不是(28, 28, 3)。resize必须和训练时一致,训练集是 28×28,你喂一张 300×300 进去,模型会直接懵掉。invert参数是关键,因为 MNIST 训练集是黑底白字,你自己的图片多数是白底黑字,不反色的话预测结果基本随机。reshape 成(1, 28, 28, 1),第一个 1 是 batch 维,模型一次处理一张。最后把 probs 里最大的索引取出来,同时返回置信度,方便在界面显示“识别为 7,置信度 98.3%”。
如果要批量测试一个文件夹里的图片,可以套一层循环,把结果写成 CSV 或直接打印。这个脚本改造成本很低,也是报告里“测试结果”部分的素材来源。
4.3 完整代码.py 和主要功能.py:什么时候用哪个
完整代码.py 的设计目的是“一键跑完”。它通常把训练、保存、加载、识别、甚至登录界面都串在同一个文件里,适合答辩现场演示。主要功能.py 则是把图像预处理、模型预测这些动作抽成函数,供其他模块复用。用的时候注意两个文件的职责区别:调试阶段不要依赖完整代码.py,它会从头训练一遍,浪费时间;应该直接用数字图片识别.py 加载已有模型,快速验证。
一个常见翻车现场是:完整代码.py 里同时导入了主要功能.py 和登录主界面.py,形成循环 import。比如登录主界面.pyimport complete_code,而完整代码.py 开头又from 登录主界面 import check_login,Python 解释器在处理循环导入时会报属性不存在。解决办法是不要在两个模块之间互相 import,统一让登录界面作为入口,内部再去调用主功能模块。这个坑在课程设计项目里非常常见,尤其是文件名带中文时,模块名解析更要注意。
如果你想快速生成一张预测效果图,可以用 PIL 在识别后显示图片和结果:
from PIL import Image import matplotlib.pyplot as plt img = Image.open("my_digit.png").convert("L") plt.imshow(img, cmap="gray") plt.title(f"pred: {pred}, conf: {conf:.2f}") plt.axis("off") plt.show()convert("L")把图片转成灰度,cmap="gray"让 matplotlib 按灰度显示。这张截图放到实验报告里,比单纯写一句“准确率 99%”更有说服力。整个链路的文件关系就是:登录主界面.py 负责开门,主要功能.py 提供服务,数字图片识别.py 做预测,完整代码.py 负责串场。你按这个思路去读代码,结构会非常清楚。
5. 避坑指南:常见问题与排查记录
这部分是从课程设计答辩现场和实际跑代码中整理出来的高频问题。每一条我都按“现象 → 原因 → 解决”的顺序说,你照着排查就行。
5.1 训练时 loss 不降或直接报错
现象:训练脚本一跑,loss 直接是 nan,或者前几个 epoch 卡在 2.3 左右一动不动。
原因:最常见的是输入没有归一化。像素值 0 到 255 直接喂进网络,卷积输出范围被放大,梯度一下爆炸。其次是标签和损失函数不匹配,比如用categorical_crossentropy但传入的是整数标签,不是 one-hot。
解决:检查 x_train 有没有除以 255;再确认 y_train 是否经过to_categorical;最后看第一个 Conv2D 的input_shape和 reshape 后数据是不是都是(28, 28, 1)。这三个查完,99% 的问题都能定位。如果 loss 还是不稳,就把学习率从默认值再降一个数量级,但一般情况下不要动这个参数。
5.2 图片识别准确率低,问题可能不在模型
现象:test 集准确率有 99%,但是把自己用手机拍的数字放进去,预测结果完全乱来。
原因:MNIST 训练集是 28×28 灰度图,黑底白字。你拍的图片是白底黑字,加上背景复杂,模型没见过这种分布。
解决:识别前先转灰度、resize 到 28×28,再做一次反色,把白底黑字变成黑底白字。我自己的经验是,加了img = 255 - img之后,准确率从 30% 直接拉到 90% 以上。这不仅是一个代码问题,也是答辩时老师最爱问的“泛化能力”问题,值得写进报告。如果反色后还是错,就检查图片比例,数字不要太小,四周留白要均匀。
5.3 运行完整代码.py 报缺模块
现象:在另一台电脑上解压项目,运行完整代码.py,报ModuleNotFoundError: No module named 'tensorflow'。
原因:代码本身没坏,是这台机器没有安装依赖。很多同学换了电脑就忘了 requirements.txt。
解决:在项目根目录执行pip install -r requirements.txt。如果安装时报 building wheel 错误,多半是 Python 版本太新,换 Python 3.8/3.9 再试。这里也推荐把pip install和python --version的输出截图留档,报告环境部分能直接用。如果你在终端里 import tensorflow 成功,但 PyCharm 里报错,那是解释器选错了,设置里把项目解释器换成 venv 里的 python.exe 就行。
5.4 数据集路径与 mnist_pic.zip 的处理
现象:运行训练脚本时报FileNotFoundError: [Errno 2] ... mnist_pic,或者训练用的数据和自己准备的图片对不上。
原因:mnist_pic.zip 还躺在压缩包里没解压,或者代码用相对路径,但你从别的目录启动脚本。更隐蔽的是,代码内部优先调用mnist.load_data(),你准备的图片集根本没被用到。
解决:先解压 mnist_pic.zip,确认目录层级和代码里写的一致。然后看代码里有没有mnist.load_data()这句,如果有,而你希望用本地图片集,把加载方式改成读文件夹。路径最好用os.path.join(os.path.dirname(__file__), 'mnist_pic'),这样在哪个目录启动都不怕。文件名里带中文也可能引发编码问题,项目根目录尽量不要放桌面上,纯英文路径最稳。
5.5 报告文档和代码版本不一致怎么办
现象:实验报告里写测试准确率 99.2%,你自己跑出来是 98.6%,担心交上去对不上。
原因:深度学习本身有随机性,不同 TensorFlow 版本、CPU/GPU、随机种子都会带来零点几个百分点的波动。这是正常现象,不是代码造假。
解决:不要改报告数字去硬凑。提交前把自己跑的训练曲线、测试准确率截图替换进报告,保持一致性。课程设计答辩更看重的是思路和能不能复现,98% 和 99% 的差距不会扣分。我在另一个项目里就吃过亏,报告里写了精确到小数点后三位的准确率,结果现场换了台电脑数值不一样,老师直接追问“是不是改过实验数据”,那场面非常被动。从那以后我交报告前一定重新跑一遍,用实际结果更新数字。
6. 进阶用法:把这份项目改成自己的课程设计
如果你不想完全照搬,下面三个方向能让你快速把它变成自己的东西,同时不被老师质疑抄袭。
6.1 改造输入:从 MNIST 到自己的手写图片
最简单的差异化是换数据来源。你可以收集 200 张自己写下的数字,按 0 到 9 分文件夹,用 3.1 里的load_from_folder()加载。为了让模型更健壮,可以做一点数据增强:
from scipy.ndimage import rotate import random angle = random.uniform(-15, 15) img_aug = rotate(img, angle, reshape=False)rotate的reshape=False保证旋转后图片尺寸不变,不会因为边缘裁切导致数字移位过多。数据量不够时,每张图旋转几次再训练,准确率通常能稳住。这个改动写进报告就是“自己采集数据并做了数据增强”,比单纯跑 MNIST 有说服力。
6.2 调参方向:卷积核数量、Dropout、Epoch
| 参数 | 建议 | 原因 |
|---|---|---|
| 卷积核数量 | 32→64→128 | 逐层增加,特征从局部到抽象 |
| Dropout | 0.25~0.5 | 太大容易欠拟合,太小防不住过拟合 |
| epochs | 10~15 | MNIST 上 15 轮后验证集收益很低 |
调参时不要一次改多个。先固定 Dropout 和 epochs,把卷积核从 32/64 改成 64/128,看 val_acc 是升是降;再单独调 Dropout。这样报告里能写清楚“控制变量法”,老师最喜欢看见这个。
6.3 验证方法:用混淆矩阵卡住报告数据
准确率是单指标,混淆矩阵能展示模型在哪个数字上容易犯错,这是报告里很加分的图:
from sklearn.metrics import confusion_matrix, classification_report import seaborn as sns y_pred = np.argmax(model.predict(x_test, verbose=0), axis=1) cm = confusion_matrix(np.argmax(y_test, axis=1), y_pred) sns.heatmap(cm, annot=True, fmt="d") plt.title("Confusion Matrix") plt.savefig("confusion_matrix.png", dpi=150) print(classification_report(np.argmax(y_test, axis=1), y_pred))np.argmax把 one-hot 标签还原成数字,classification_report会输出每一类的精确率、召回率和 F1 值。你把这些数据写进实验报告,再配合矩阵图,老师一眼就能看出你做了完整验证。如果发现 4 和 9 经常混,说明训练样本里这两个数字的风格太接近,可以针对性补几张磨砂边缘的样本。
从那以后,我每次拿到别人分享的课程设计项目,都强制自己先走一遍最小闭环:装环境、训练、预测,再回去看文档。凡是省略这一步的,最后无一例外都在答辩现场翻车。希望这篇笔记帮到你。
本文还有配套的精品资源,点击获取