news 2026/9/26 11:31:41

CNN手写数字识别可视化:从卷积到分类概率的完整演示

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
CNN手写数字识别可视化:从卷积到分类概率的完整演示

卷积神经网络居然能这么讲?手写数字识别可视化,一条视频讲透 CNN

这次我们来看一个特别的“项目”:用可视化方式,把卷积神经网络(CNN)识别手写数字的完整过程,1 分钟之内掰开揉碎讲清楚。

很多人学 CNN 卡住,不是因为数学看不懂,而是因为看不到网络内部到底在干什么。输入一张数字图,经过卷积、池化、全连接,每一步输出是什么形状、特征图长什么样、数字是怎么被“认”出来的,全凭想象。可视化演示解决的就是这个问题:让卷积核的滑动、特征图的逐层变化、最后的分类概率全部直接画在屏幕上,看完你就能在脑子里建立起 CNN 的完整运行画面。

这篇文章既是对这类可视化项目的拆解,也是给想自己动手复现“手写数字识别可视化”的读者的一份实操指南。我会讲清楚 CNN 核心模块怎么理解,可视化要展示哪些关键节点,代码和数据集怎么准备,以及部署和调试时最容易踩的坑。无论你是初学者想搞懂 CNN,还是想给课程、博客、视频做一个能跑的可视化 Demo,这篇都值得直接收藏。

1. CNN 手写数字识别可视化:核心能力速览

先给规格,再讲细节。这个主题包含的是一个典型的“深度学习模型训练 + 可视化展示”的完整流程,整体能力可以按下表理解。

能力项说明
核心任务手写数字识别(MNIST / 自绘数字图片分类)
基础网络卷积层(Conv2d)+ 汇聚层(Pooling)+ 全连接层(FC)
可视化内容卷积核滑动过程、特征图输出、池化降维效果、分类概率分布
视觉呈现逐层特征图打印、单张数字预测过程、1 分钟短视频式演示
部署方式Python 脚本 / Jupyter Notebook / 简单 Web 页面
API 能力可通过 Flask/FastAPI 封装预测接口(按需扩展)
批量任务支持对测试集批量识别并统计准确率
适合场景CNN 入门教学、模型可视化讲解、课程作业、技术博客配图

从材料看,项目的主体是“教学演示 + 可视化”,不是高门槛 AI 产品。CPU 就能跑,不需要独立显卡,不需要大显存,MNIST 数据集也非常轻量。读者要准备的无非是一个 Python 环境、PyTorch 或 TensorFlow、以及一个能显示图片的 Notebook 环境。

2. 为什么可视化对手写数字识别这么重要

CNN 的入门门槛不在“会调用模型”,而在“理解模型内部发生了什么”。手写数字识别是最常用的教学任务,原因是数据集简单、类别固定、训练快,但它同样包含卷积、池化、全连接、激活、损失计算、概率输出这些完整流程。把这些流程可视化以后,至少能解决三个学习痛点。

第一,卷积层不再抽象。很多人背住了“卷积核提取特征”这句话,但不知道 3x3 的卷积核在 28x28 的图像上怎么滑动、每个位置和图像块做了什么运算。可视化可以直接画出滑动过程,每一步选中一个局部区域,做逐元素相乘再求和,生成一个新的像素值。看懂这一段,CNN 就不再是黑盒。

第二,特征图的“逐层变化”能看见。第一层卷积输出往往还保留数字的轮廓;第二层卷积开始提取边缘、角落、笔画交叉等更抽象的结构;池化层把分辨率降下来,特征图变小但信息更集中。可视化把这些特征图按层平铺打印,一眼就能看出网络在“看”什么。

第三,最后的分类概率最有说服力。输入一个“7”,网络输出的是一个 10 维向量,每个值代表属于 0 到 9 的概率。可视化把概率画成柱状图,看到 7 的概率接近 1、其他接近 0,整个模型的推理逻辑就闭环了。

手写数字识别可视化的核心价值,就是把“输入 -> 卷积 -> 池化 -> 全连接 -> 输出概率”这条链路变成肉眼可见的过程。理解了这个链路,再去看 LeNet-5、ResNet、YOLO 这些复杂网络,思路会顺很多。

3. 环境准备与前置条件

这类可视化项目对环境要求很低。下面给出一套通用检查清单,版本不需要完全照抄,只要符合自己本机情况即可。

3.1 硬件门槛

  • CPU 推理完全够用,MNIST 单张 28x28 灰度图,训练一轮也很快。
  • 可选 NVIDIA GPU 加速,但最低端显卡都能跑这个任务。
  • 内存建议 4GB 以上,磁盘预留 2GB 左右装环境和模型文件。

3.2 软件依赖

推荐 Python 3.8 及以上,安装以下核心库:

pip install torch torchvision matplotlib numpy

如果希望用更简单的接口快速演示,可以加装:

pip install jupyter notebook

如果计划把可视化做成 Web 页面或 API 服务,再安装:

pip install flask

3.3 数据集准备

MNIST 是手写数字识别的经典数据集:

  • 训练集 60000 张,测试集 10000 张。
  • 图像尺寸 28x28,单通道灰度图。
  • 类别为 0 到 9 共 10 个数字。

使用 PyTorch 时,直接通过 torchvision 下载即可:

from torchvision import datasets, transforms transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_data = datasets.MNIST(root='./data', train=True, download=True, transform=transform) test_data = datasets.MNIST(root='./data', train=False, download=True, transform=transform)

注意:国内网络环境下,如果 torchvision 下载 MNIST 失败,可以先手动下载数据集文件,放到./data/MNIST/raw目录下。

4. CNN 结构设计与可视化节点规划

可视化项目要能讲清楚,网络结构不能太复杂。这里采用最经典的 LeNet-5 简化结构,既保留 CNN 核心模块,又方便逐层展示。

网络结构如下:

层名称类型输出形状作用
输入层原始图像(1, 28, 28)灰度数字图
Conv1卷积层(32, 26, 26)提取边缘、纹理等低级特征
Pool1汇聚层(32, 13, 13)下采样,缩小特征图尺寸
Conv2卷积层(64, 11, 11)提取笔画、结构等抽象特征
Pool2汇聚层(64, 5, 5)进一步降维
Flatten展平1600转成一维向量
FC1全连接层128特征组合
FC2全连接层10输出 10 个类别得分

对应代码如下:

import torch.nn as nn class CNNNet(nn.Module): def __init__(self): super().__init__() self.conv1 = nn.Conv2d(1, 32, kernel_size=3) self.pool1 = nn.MaxPool2d(2) self.conv2 = nn.Conv2d(32, 64, kernel_size=3) self.pool2 = nn.MaxPool2d(2) self.fc1 = nn.Linear(1600, 128) self.fc2 = nn.Linear(128, 10) def forward(self, x): x = self.pool1(torch.relu(self.conv1(x))) x = self.pool2(torch.relu(self.conv2(x))) x = x.view(x.size(0), -1) x = torch.relu(self.fc1(x)) x = self.fc2(x) return x

可视化规划时,建议展示以下节点:

  • 输入图像:展示原始手写数字。
  • 第一个卷积层输出:把 32 张特征图用网格方式画出来。
  • 第一个池化层输出:对比池化前后特征图的尺寸变化。
  • 第二个卷积层输出:挑选若干特征图展示。
  • 展平和全连接阶段:可以画出向量维度的变化。
  • 最终输出:绘制概率柱状图。

一个关键技巧是:hook 网络中间层,在 forward 过程中把每一层的输出抓出来,不需要修改网络结构就能可视化。示例:

activations = {} def hook_fn(name): def fn(model, input, output): activations[name] = output.detach() return fn model.conv1.register_forward_hook(hook_fn('conv1')) model.pool1.register_forward_hook(hook_fn('pool1')) model.conv2.register_forward_hook(hook_fn('conv2')) model.pool2.register_forward_hook(hook_fn('pool2'))

这样在模型推理结束后,中间层输出会保存到字典里,后面统一画图。

5. 训练流程与模型保存

可视化演示需要一个训练好的模型。训练代码不用太复杂,MNIST 任务在 CPU 上也能十几分钟跑完,有 GPU 更快。

下面给出一套最小可运行训练脚本:

import torch from torch import nn, optim from torch.utils.data import DataLoader from torchvision import datasets, transforms transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_data = datasets.MNIST(root='./data', train=True, download=True, transform=transform) test_data = datasets.MNIST(root='./data', train=False, download=True, transform=transform) train_loader = DataLoader(train_data, batch_size=64, shuffle=True) test_loader = DataLoader(test_data, batch_size=64, shuffle=False) model = CNNNet() criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=0.001) for epoch in range(5): running_loss = 0.0 for images, labels in train_loader: optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() running_loss += loss.item() print(f"epoch {epoch+1}, loss: {running_loss/len(train_loader):.4f}") # 测试集准确率 correct = 0 total = 0 model.eval() with torch.no_grad(): for images, labels in test_loader: outputs = model(images) _, predicted = torch.max(outputs.data, 1) total += labels.size(0) correct += (predicted == labels).sum().item() print(f"test accuracy: {100 * correct / total:.2f}%") torch.save(model.state_dict(), "./mnist_cnn.pth")

训练完成后,mnist_cnn.pth就是用来做可视化和单张预测的权重文件。后续演示只需要加载这个文件,不需要每次重新训练。

6. 可视化演示:单张数字识别全过程

这是整个项目的核心看点。加载一个训练好的模型,输入一张手写数字图,依次展示卷积滑动、特征图变化和最终分类结果。

6.1 加载模型和输入图片

import torch import matplotlib.pyplot as plt from PIL import Image import torchvision.transforms as transforms model = CNNNet() model.load_state_dict(torch.load("./mnist_cnn.pth", map_location="cpu")) model.eval() # 读取一张 28x28 的灰度手写数字图 img = Image.open("./test_digit.png").convert("L").resize((28, 28)) transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) input_tensor = transform(img).unsqueeze(0)

6.2 展示单张图片

plt.figure(figsize=(3, 3)) plt.imshow(img, cmap="gray") plt.title("input digit") plt.axis("off") plt.show()

6.3 绘制各层特征图

利用前面注册的 hook,模型推理一次后,把每一层输出画成网格图:

with torch.no_grad(): output = model(input_tensor) def show_activation_grid(activation, title): activation = activation.squeeze(0) num_channels = activation.shape[0] cols = 8 rows = (num_channels + cols - 1) // cols fig, axes = plt.subplots(rows, cols, figsize=(cols * 1.5, rows * 1.5)) for i in range(num_channels): ax = axes[i // cols, i % cols] ax.imshow(activation[i].cpu().numpy(), cmap="viridis") ax.axis("off") plt.suptitle(title) plt.show() show_activation_grid(activations['conv1'], "conv1 feature maps") show_activation_grid(activations['pool1'], "pool1 feature maps") show_activation_grid(activations['conv2'], "conv2 feature maps") show_activation_grid(activations['pool2'], "pool2 feature maps")

从实际展示效果看,conv1 的特征图能明显看出数字的亮色轮廓,conv2 的特征图开始出现笔画强化和背景抑制的效果,pool 层的特征图分辨率变小但结构轮廓仍然清楚。这种逐层变化正是可视化教学最直观的部分。

6.4 画出分类概率柱状图

import numpy as np probs = torch.softmax(output, dim=1).squeeze(0).numpy() plt.figure(figsize=(8, 4)) plt.bar(range(10), probs) plt.xticks(range(10)) plt.title("classification probability") plt.xlabel("digit class") plt.ylabel("probability") plt.show()

最终预测结果可以直接取概率最大的索引:

pred = torch.argmax(output, dim=1).item() print(f"predict: {pred}")

这一步完成后,“一张图输入 -> 逐层特征可视化 -> 概率输出”的完整链路就对用户完全可见了。

7. 功能测试与效果验证

可视化项目也要走“先测基础,再测特殊样本”的流程,否则很容易出现“训练准确率很高、但演示用例全错”的情况。

7.1 标准测试集验证

这是判断模型本身是否可靠的底线指标:

total_correct = 0 total_count = 0 with torch.no_grad(): for images, labels in test_loader: outputs = model(images) _, predicted = torch.max(outputs, 1) total_count += labels.size(0) total_correct += (predicted == labels).sum().item() print(f"test accuracy: {100 * total_correct / total_count:.2f}%")

按照正常训练流程,5 个 epoch 后准确率一般能到 98% 以上。如果低于 90%,优先检查数据归一化参数、学习率和网络结构。

7.2 人工绘制数字识别

这是可视化演示里最有说服力的一环:

  1. 用画图软件或 PIL 手动画一个数字,保存为 png。
  2. 转换为灰度图并 resize 到 28x28。
  3. 用模型做预测并画出特征图。

测试时要特别关注手写风格较随意的数字,比如断笔的 4、写得像 7 的 1、带斜线的 0。这些边界样本最能暴露模型和可视化的判断差异。

7.3 错误样本分析

从测试集里筛选预测错误的样本,打印输入图和真实标签,再看特征图变化,通常能发现两类问题:

  • 数字潦草到人类也难以分辨,模型犯错合理。
  • 训练不充分或过拟合导致简单样本判断失误,此时需要调整训练策略。

7.4 判断可视化是否成功

一个合格的 CNN 手写数字识别可视化,必须呈现以下结果:

  • 输入图像显示正常。
  • conv1 和 pool1 的特征图尺寸符合预期,图像内容仍可见。
  • conv2 的特征图出现更抽象的结构响应。
  • 最终概率柱状图有明确主导类别。
  • 预测结果与实际数字一致,或错误样本能给出合理解释。

如果特征图全黑或全白,大概率是模型未收敛或数据归一化错误;如果概率值均匀分布,说明模型对当前输入没有足够判断依据。

8. 接口 API 与批量预测扩展

演示项目如果只停在 Notebook 里,作用有限。更实用的做法是把它封装成 API 服务,或者做批量预测,这样就能接入其他工具。

8.1 Flask 预测接口

下面是一个最小可用的预测接口示例:

import io import torch from flask import Flask, request, jsonify from PIL import Image import torchvision.transforms as transforms import torch.nn as nn class CNNNet(nn.Module): # 与训练时的网络结构保持一致 pass app = Flask(__name__) model = CNNNet() model.load_state_dict(torch.load("./mnist_cnn.pth", map_location="cpu")) model.eval() transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) @app.route("/predict", methods=["POST"]) def predict(): file = request.files["file"] img = Image.open(io.BytesIO(file.read())).convert("L").resize((28, 28)) tensor = transform(img).unsqueeze(0) with torch.no_grad(): output = model(tensor) prob = torch.softmax(output, dim=1).squeeze(0) pred = torch.argmax(prob).item() return jsonify({"prediction": pred, "probability": prob.tolist()}) if __name__ == "__main__": app.run(host="127.0.0.1", port=5000)

启动后,可以用 curl 测试接口:

curl -X POST http://127.0.0.1:5000/predict -F "file=@./test_digit.png"

这个接口可以直接接到网页上传按钮或者自动化流程里。

8.2 Python 批量识别测试

如果要验证模型对整张多数字图片的识别效果,先做好图像切分,再逐块送入模型,最后合并结果:

from PIL import Image def batch_predict(image_path, cell_size=28): img = Image.open(image_path).convert("L") # 假设图片是水平排列的多个数字 width, height = img.size results = [] for x in range(0, width, cell_size): crop = img.crop((x, 0, x + cell_size, height)) tensor = transform(crop).unsqueeze(0) with torch.no_grad(): output = model(tensor) pred = torch.argmax(output, dim=1).item() results.append(pred) return results # 示例:识别一行手写数字字符串 print(batch_predict("./multi_digit.png"))

这里强调一下:实际切分方式要按图片布局来定,不能假设所有图片都是等宽切分。通用做法是先做连通域检测,再按连通域位置切分单个数字。

9. 资源占用与性能观察

这个项目资源占用很低,但可视化过程仍然有一些值得观察的点。

9.1 CPU 训练与推理

MNIST 单张图片推理在 CPU 上通常是毫秒级,训练一个 epoch 在普通笔记本上大约几十秒到两分钟,整体占用内存不超过 1GB。可视化阶段的主要开销是 matplotlib 绘图,特征图数量多时 CPU 会短暂升高,但不会影响其他程序。

9.2 显存占用

如果使用 GPU,显存占用会非常低,通常几十 MB 到一两百 MB,因为网络尺寸小、输入分辨率低。如果用户拿到的是一键包版本,需要注意 CUDA 版本和 PyTorch 版本的兼容性。要求不高的场景直接使用 CPU 推理反而更省事。

9.3 批量任务性能

批量预测在 GPU 上的收益明显。一次送入 64 张图片比逐张推理快很多:

# GPU 批量预测 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = model.to(device) with torch.no_grad(): for images, labels in test_loader: images, labels = images.to(device), labels.to(device) outputs = model(images)

如果显存不足,把 batch_size 调小即可;如果 CPU 内存不足,减少测试样本数量或按小批量遍历。

10. 常见问题与排查方法

问题现象可能原因排查方式解决方案
数据集下载失败torchvision 自动下载被网络限制查看脚本报错信息手动下载 MNIST 文件放入 raw 目录
训练 loss 不下降学习率过高或网络结构错误打印前几轮 loss 变化降低学习率,检查 forward 流程
特征图全黑模型未收敛或归一化参数错误打印特征图数值分布检查 Normalize 参数和训练是否完成
概率分布均匀输入图片预处理错误输出 tensor 形状和数值确认 resize、灰度化和归一化一致
预测结果总是一个固定数字模型权重未加载或模型结构不匹配检查权重加载报错初始化网络后再 load_state_dict
Flask 接口返回 400请求缺少 file 字段或文件损坏打印 request.files用 curl 检查字段名是否正确
批量识别切分错误图片布局不均匀打印切分后的图片改用连通域检测或人工标注切分位置
显卡可用但 PyTorch 不识别CUDA 版本不匹配运行 torch.cuda.is_available()安装匹配 CUDA 版本的 PyTorch
可视化页面卡顿matplotlib 绘制特征图太多观察 CPU 占用减少展示通道数,只选 8 张代表性特征图

重点排查思路是:先确认数据预处理一致,再确认模型能跑通单张推理,最后再做批量或接口扩展。可视化异常大多数不是绘图代码出错,而是上游数据或模型输出本身有问题。

11. 最佳实践与使用建议

如果你准备自己写一个 CNN 手写数字识别可视化项目,或者打算把这个内容做成课程、博客、视频演示,下面几条建议可以直接拿来用。

第一,先跑通再美化。第一版代码用 Mini 数据集、少通道数、只画 4 张特征图,确认链路通了,再逐步增加特征图展示和动画效果。一上来就做 1 分钟动画,调试成本会很高。

第二,统一图像预处理。训练、测试、可视化、API 四个环节的 transform 必须保持一致。很多“模型预测不准”的问题,本质上是对输入图片做了不同的 resize 或归一化。

第三,保存一份固定随机种子。为了训练结果可复现,在训练前设置随机种子:

import random import numpy as np import torch random.seed(42) np.random.seed(42) torch.manual_seed(42)

这样每次训练结果基本一致,可视化演示时不会出现“上次准确率 99%,这次只有 90%”的尴尬。

第四,把中间层输出保存成图片。不要只在 Jupyter 里展示,可以把特征图保存到本地,方便后续插入博客、PPT 或者视频:

import matplotlib matplotlib.use("Agg") for i in range(8): plt.subplot(2, 4, i + 1) plt.imshow(activations['conv1'][0][i].cpu().numpy(), cmap="viridis") plt.axis("off") plt.savefig("./feature_map_conv1.png", dpi=150)

第五,注意边界样本。手写数字识别的鲁棒性测试,建议加入旋转、平移、粗细变化、噪声干扰等样本,可视化能直观暴露模型的薄弱点。

第六,如果做视频化演示,建议每一层停留 3 到 5 秒,标注张量维度的变化,比如“1x28x28 -> 32x26x26 -> 32x13x13”。人对维度变化的理解比对像素值的理解更慢,标注能大大降低认知负担。

12. 总结与下一步

卷积神经网络可视化最值得做的一件事,就是让学习者不再靠背诵公式来理解 CNN,而是直接看到每一个卷积核、每一层池化、每一次概率输出。手写数字识别是这个目标的最佳载体:任务简单、数据轻量、效果直观,而且从训练到可视化到 API 部署,整套流程一个人在一台普通笔记本上就能完成。

如果你是由零开始,我建议按这个顺序来验证:

  1. 跑通训练脚本,确认测试集准确率达到 97% 以上。
  2. 加载单张图片,画出 conv1、pool1、conv2、pool2 的特征图。
  3. 打印最终概率柱状图,确认分类结果。
  4. 用 Flask 封装一个预测接口,测试 curl 调用。
  5. 做一批困难样本测试,记录模型误判情况。

最容易踩的坑集中在三处:数据预处理不统一、模型权重加载失败后的静默错误、特征图绘制时张量维度理解错误。把这三个问题提前解决,这个项目基本不会再卡人。

后续可以扩展的方向不少:把手写数字识别改成树叶分类、票据文字识别、验证码识别;把静态特征图改成动态动画,用 matplotlib animation 把卷积核滑动过程做成 gif;在 API 接口上加入批量识别和日志记录,接入完整的数据标注、训练、部署链路;甚至可以把可视化 Web 化,做成浏览器里可交互的 CNN 演示页面。每一个方向都能进一步加深对卷积神经网络的理解,而且都可以从今天这个手写数字识别可视化项目里直接延伸。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/26 11:31:18

PyTorch实战:MNIST手写数字识别从训练到部署全流程

简介:这是一份面向Python初学者与深度学习入门者的手写数字识别实战资源,围绕卷积神经网络(CNN)识别手写数字这一经典计算机视觉任务展开,帮助读者理解图像特征提取与分类的完整流程。压缩包共13个文件,约6…

作者头像 李华
网站建设 2026/9/26 11:31:15

UE项目如何接入大语言模型:Qwen3.8 27B本地部署与云端模型选择指南

如果你的日常开发工作已经离不了大语言模型,那么最近这段时间你一定会陷入一种“选择困难”:本地能跑的模型越来越强,云端 API 的版本迭代越来越快,而你的实际场景又往往是“要在一个具体的引擎或工具链里把模型用起来”&#xff…

作者头像 李华
网站建设 2026/9/26 11:30:43

GCC 9.3.0源码编译实战:从解压到安装避坑指南

简介:GCC 9.3.0 源码包面向需要指定编译器版本进行环境构建、源码阅读或二次开发的工程师,以及在 RHEL/CentOS 7 等老旧系统中替换或扩展系统自带 GCC 的场景,可直接离线获取,免去在线检索与下载的不确定性。压缩包约 118.39MB&am…

作者头像 李华
网站建设 2026/9/26 11:30:41

星辰变归来正版官方客户端下载指引,忆往游戏正规安全渠道指南

《星辰变归来》由安徽游昕网络科技有限公司联合忆往游戏平台负责运营,是经过正版授权、改编自经典网文 IP《星辰变》的 3D 修真怀旧手游。现阶段游戏依托专属官方主站面向全网正式开放,高度复刻经典端游原版修真体系与世界观,坚持绿色公平长久…

作者头像 李华
网站建设 2026/9/26 11:27:58

MediaPipe+SVM手势数字识别实战:端到端可复现机器学习pipeline

简介:本资源是一个基于MediaPipe的手势数字识别机器学习实战项目,面向计算机、人工智能、数据科学等专业学生及初入AI领域的开发者,解决手势图像采集、关键点提取、数字分类建模与实时识别等核心问题,适用于课程设计、大作业或入门…

作者头像 李华