MNIST 手写数字识别实战:CNN + Transformer 混合模型

文章目录

  • [MNIST 手写数字识别实战:CNN + Transformer 混合模型](#MNIST 手写数字识别实战:CNN + Transformer 混合模型)

MNIST 手写数字识别实战:CNN + Transformer 混合模型

一、项目简介

本项目用 PyTorch 在 CPU 上 训练一个手写数字识别模型。与常见的纯 CNN 方案不同,这里采用了混合架构:CNN 负责提取局部笔画,Transformer 负责建模全局关系。

模型约 95k 参数,CPU 训练 10--15 分钟 即可在官方测试集上达到 98%+ 准确率。项目注释全部是中文、面向初学者,还配套了多种可视化手段,非常适合入门深度学习与注意力机制。

打个比方:识别手写数字,就像一群人一起猜「这张纸条上写的是几」。CNN 是拿着放大镜的人,只看局部;Transformer 是开讨论会的人,让每一块拼图互相看看、彼此印证。

整条流水线一图看懂:
#mermaid-svg-gEE6XBMk84u3SwI0{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-gEE6XBMk84u3SwI0 .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-gEE6XBMk84u3SwI0 .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-gEE6XBMk84u3SwI0 .error-icon{fill:#552222;}#mermaid-svg-gEE6XBMk84u3SwI0 .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-gEE6XBMk84u3SwI0 .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-gEE6XBMk84u3SwI0 .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-gEE6XBMk84u3SwI0 .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-gEE6XBMk84u3SwI0 .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-gEE6XBMk84u3SwI0 .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-gEE6XBMk84u3SwI0 .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-gEE6XBMk84u3SwI0 .marker{fill:#333333;stroke:#333333;}#mermaid-svg-gEE6XBMk84u3SwI0 .marker.cross{stroke:#333333;}#mermaid-svg-gEE6XBMk84u3SwI0 svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-gEE6XBMk84u3SwI0 p{margin:0;}#mermaid-svg-gEE6XBMk84u3SwI0 .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-gEE6XBMk84u3SwI0 .cluster-label text{fill:#333;}#mermaid-svg-gEE6XBMk84u3SwI0 .cluster-label span{color:#333;}#mermaid-svg-gEE6XBMk84u3SwI0 .cluster-label span p{background-color:transparent;}#mermaid-svg-gEE6XBMk84u3SwI0 .label text,#mermaid-svg-gEE6XBMk84u3SwI0 span{fill:#333;color:#333;}#mermaid-svg-gEE6XBMk84u3SwI0 .node rect,#mermaid-svg-gEE6XBMk84u3SwI0 .node circle,#mermaid-svg-gEE6XBMk84u3SwI0 .node ellipse,#mermaid-svg-gEE6XBMk84u3SwI0 .node polygon,#mermaid-svg-gEE6XBMk84u3SwI0 .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-gEE6XBMk84u3SwI0 .rough-node .label text,#mermaid-svg-gEE6XBMk84u3SwI0 .node .label text,#mermaid-svg-gEE6XBMk84u3SwI0 .image-shape .label,#mermaid-svg-gEE6XBMk84u3SwI0 .icon-shape .label{text-anchor:middle;}#mermaid-svg-gEE6XBMk84u3SwI0 .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-gEE6XBMk84u3SwI0 .rough-node .label,#mermaid-svg-gEE6XBMk84u3SwI0 .node .label,#mermaid-svg-gEE6XBMk84u3SwI0 .image-shape .label,#mermaid-svg-gEE6XBMk84u3SwI0 .icon-shape .label{text-align:center;}#mermaid-svg-gEE6XBMk84u3SwI0 .node.clickable{cursor:pointer;}#mermaid-svg-gEE6XBMk84u3SwI0 .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-gEE6XBMk84u3SwI0 .arrowheadPath{fill:#333333;}#mermaid-svg-gEE6XBMk84u3SwI0 .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-gEE6XBMk84u3SwI0 .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-gEE6XBMk84u3SwI0 .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-gEE6XBMk84u3SwI0 .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-gEE6XBMk84u3SwI0 .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-gEE6XBMk84u3SwI0 .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-gEE6XBMk84u3SwI0 .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-gEE6XBMk84u3SwI0 .cluster text{fill:#333;}#mermaid-svg-gEE6XBMk84u3SwI0 .cluster span{color:#333;}#mermaid-svg-gEE6XBMk84u3SwI0 div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-gEE6XBMk84u3SwI0 .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-gEE6XBMk84u3SwI0 rect.text{fill:none;stroke-width:0;}#mermaid-svg-gEE6XBMk84u3SwI0 .icon-shape,#mermaid-svg-gEE6XBMk84u3SwI0 .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-gEE6XBMk84u3SwI0 .icon-shape p,#mermaid-svg-gEE6XBMk84u3SwI0 .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-gEE6XBMk84u3SwI0 .icon-shape .label rect,#mermaid-svg-gEE6XBMk84u3SwI0 .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-gEE6XBMk84u3SwI0 .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-gEE6XBMk84u3SwI0 .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-gEE6XBMk84u3SwI0 :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 输入灰度图 B, 1, 28, 28
CNN:两层 Conv+BN+ReLU+MaxPool

放大镜扫出局部笔画

B, 64, 7, 7
proj: 1×1 卷积改写成

Transformer 认识的格式
展平成 49 个 token

  • 请来班长 CLS、发座位卡 pos_embed

B, 50, 64
Transformer Encoder × 2 层

开讨论会建模全局关系
取班长发言 tokens:,0

经 LayerNorm B, 64
Linear(64→10)

压成 10 个分数 logits

二、CNN:提取局部笔画

CNN 部分就是模型里的 self.cnn,由两个相同的「卷积块」组成,每个块包含 4 层:

复制代码
Conv2d → BatchNorm2d → ReLU → MaxPool2d

1. 卷积核:3×3 的放大镜

卷积用 3×3 的小窗口在 28×28 的图上滑动,每滑到一处就问「这里像某种笔画吗?」。padding=1 保证滑动过程中图不缩水,只有池化才会缩小。

2. 通道:64 支荧光笔

  • **通道(channel)**可以想成多种荧光笔:第一层用 32 支 笔标 32 种局部模式(竖线、横线、弧线...),第二层再用 64 支笔标更复杂的组合(比如「竖线 + 小圈」)。
  • 每支笔 = 一个卷积核,训练就是让笔学会"认出"自己的那种模式。

3. BatchNorm 与 ReLU

  • BatchNorm:把每支笔的深浅拉到差不多的尺度,训练更稳;
  • ReLU:负数打成 0------没看到这种笔画就记 0,看到了才往上报。

4. MaxPool:把地图折小

MaxPool2d(2) 在每 2×2 的格子里只保留最大值,图缩小一半:28 → 14 → 7。细节少了,但「大概长什么样」还在,计算量也大大降低。

最终得到 7×7 的特征图 = 49 个位置,每个位置都是 64 维的特征向量。

三、Transformer:建模全局关系

CNN 的放大镜只看局部,但「左边一竖 + 右边一个圈」拼起来才是「9」。这一步交给 Transformer。

1. 把特征图变成一串 token

  • self.proj:一个 1×1 卷积,在同一个像素位置把 64 个通道混成 d_model=64 维,相当于把特征改写为 Transformer 认识的格式;
  • flatten + transpose:把 7×7 的格子拉成一排,得到 49 个 token,每个 token 是一个 64 维向量。

2. CLS token:请一位班长

拼接一个虚拟的 CLS tokencls_token,代码里戏称「班长」),图上并不存在这一块。训练时它会慢慢学会:听完所有小块的发言,给出整张图的总结向量。序列长度变成 50

3. 位置编码:给每个座位发座位卡

Transformer 默认把 50 个向量看成「一袋没有顺序的卡片」,不加位置就分不清「左上角的圈」和「右下角的圈」。所以给每个位置加一个可学习的 pos_embed(座位卡,也是 64 维向量)。

4. 自注意力:开讨论会

nn.TransformerEncoder 由 2 层 TransformerEncoderLayer 组成,每层做两件事:

  1. 多头自注意力(开会) :每个 token 都去看其他所有 token,算出该更关注谁。nhead=4 表示 4 个小组同时讨论------有的看「是不是闭合圆」,有的看「有没有竖线」,最后把意见汇总;
  2. 前馈网络(消化)dim_feedforward=128,每个人听完会后到里屋用更宽的草稿纸把想法整理一遍。

其中 norm_first=True 先把向量标准化再计算,像考试前先把分数换成标准分;2 层就是「开两轮会」,第二层在第一层讨论结果上继续。

注意力本身怎么算?每个 token 举着 Q 提问、亮出 K 名片、掏出 V 内容,一轮下来就完成「按相关程度加权听讲」:
#mermaid-svg-WoC3AElc3Tt5Xucl{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-WoC3AElc3Tt5Xucl .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-WoC3AElc3Tt5Xucl .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-WoC3AElc3Tt5Xucl .error-icon{fill:#552222;}#mermaid-svg-WoC3AElc3Tt5Xucl .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-WoC3AElc3Tt5Xucl .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-WoC3AElc3Tt5Xucl .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-WoC3AElc3Tt5Xucl .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-WoC3AElc3Tt5Xucl .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-WoC3AElc3Tt5Xucl .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-WoC3AElc3Tt5Xucl .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-WoC3AElc3Tt5Xucl .marker{fill:#333333;stroke:#333333;}#mermaid-svg-WoC3AElc3Tt5Xucl .marker.cross{stroke:#333333;}#mermaid-svg-WoC3AElc3Tt5Xucl svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-WoC3AElc3Tt5Xucl p{margin:0;}#mermaid-svg-WoC3AElc3Tt5Xucl .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-WoC3AElc3Tt5Xucl .cluster-label text{fill:#333;}#mermaid-svg-WoC3AElc3Tt5Xucl .cluster-label span{color:#333;}#mermaid-svg-WoC3AElc3Tt5Xucl .cluster-label span p{background-color:transparent;}#mermaid-svg-WoC3AElc3Tt5Xucl .label text,#mermaid-svg-WoC3AElc3Tt5Xucl span{fill:#333;color:#333;}#mermaid-svg-WoC3AElc3Tt5Xucl .node rect,#mermaid-svg-WoC3AElc3Tt5Xucl .node circle,#mermaid-svg-WoC3AElc3Tt5Xucl .node ellipse,#mermaid-svg-WoC3AElc3Tt5Xucl .node polygon,#mermaid-svg-WoC3AElc3Tt5Xucl .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-WoC3AElc3Tt5Xucl .rough-node .label text,#mermaid-svg-WoC3AElc3Tt5Xucl .node .label text,#mermaid-svg-WoC3AElc3Tt5Xucl .image-shape .label,#mermaid-svg-WoC3AElc3Tt5Xucl .icon-shape .label{text-anchor:middle;}#mermaid-svg-WoC3AElc3Tt5Xucl .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-WoC3AElc3Tt5Xucl .rough-node .label,#mermaid-svg-WoC3AElc3Tt5Xucl .node .label,#mermaid-svg-WoC3AElc3Tt5Xucl .image-shape .label,#mermaid-svg-WoC3AElc3Tt5Xucl .icon-shape .label{text-align:center;}#mermaid-svg-WoC3AElc3Tt5Xucl .node.clickable{cursor:pointer;}#mermaid-svg-WoC3AElc3Tt5Xucl .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-WoC3AElc3Tt5Xucl .arrowheadPath{fill:#333333;}#mermaid-svg-WoC3AElc3Tt5Xucl .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-WoC3AElc3Tt5Xucl .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-WoC3AElc3Tt5Xucl .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-WoC3AElc3Tt5Xucl .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-WoC3AElc3Tt5Xucl .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-WoC3AElc3Tt5Xucl .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-WoC3AElc3Tt5Xucl .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-WoC3AElc3Tt5Xucl .cluster text{fill:#333;}#mermaid-svg-WoC3AElc3Tt5Xucl .cluster span{color:#333;}#mermaid-svg-WoC3AElc3Tt5Xucl div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-WoC3AElc3Tt5Xucl .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-WoC3AElc3Tt5Xucl rect.text{fill:none;stroke-width:0;}#mermaid-svg-WoC3AElc3Tt5Xucl .icon-shape,#mermaid-svg-WoC3AElc3Tt5Xucl .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-WoC3AElc3Tt5Xucl .icon-shape p,#mermaid-svg-WoC3AElc3Tt5Xucl .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-WoC3AElc3Tt5Xucl .icon-shape .label rect,#mermaid-svg-WoC3AElc3Tt5Xucl .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-WoC3AElc3Tt5Xucl .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-WoC3AElc3Tt5Xucl .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-WoC3AElc3Tt5Xucl :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 50 个 token

每人一份向量
Q 提问

「我是竖线的一部分,

谁和我凑成 9?」
K 名片标题

报出自己的特征类型
V 名片正文

携带真正的内容
和所有 K 匹配打分

÷√64 缩放
softmax 变权重

合计=1:这一票怎么分
加权平均
每个 token 得到新向量:

混入了最相关位置的信息

单层 Transformer 的两道工序(开会 → 消化)与残差近路:
#mermaid-svg-H6FqUY7dgiws5mwm{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-H6FqUY7dgiws5mwm .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-H6FqUY7dgiws5mwm .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-H6FqUY7dgiws5mwm .error-icon{fill:#552222;}#mermaid-svg-H6FqUY7dgiws5mwm .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-H6FqUY7dgiws5mwm .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-H6FqUY7dgiws5mwm .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-H6FqUY7dgiws5mwm .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-H6FqUY7dgiws5mwm .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-H6FqUY7dgiws5mwm .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-H6FqUY7dgiws5mwm .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-H6FqUY7dgiws5mwm .marker{fill:#333333;stroke:#333333;}#mermaid-svg-H6FqUY7dgiws5mwm .marker.cross{stroke:#333333;}#mermaid-svg-H6FqUY7dgiws5mwm svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-H6FqUY7dgiws5mwm p{margin:0;}#mermaid-svg-H6FqUY7dgiws5mwm .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-H6FqUY7dgiws5mwm .cluster-label text{fill:#333;}#mermaid-svg-H6FqUY7dgiws5mwm .cluster-label span{color:#333;}#mermaid-svg-H6FqUY7dgiws5mwm .cluster-label span p{background-color:transparent;}#mermaid-svg-H6FqUY7dgiws5mwm .label text,#mermaid-svg-H6FqUY7dgiws5mwm span{fill:#333;color:#333;}#mermaid-svg-H6FqUY7dgiws5mwm .node rect,#mermaid-svg-H6FqUY7dgiws5mwm .node circle,#mermaid-svg-H6FqUY7dgiws5mwm .node ellipse,#mermaid-svg-H6FqUY7dgiws5mwm .node polygon,#mermaid-svg-H6FqUY7dgiws5mwm .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-H6FqUY7dgiws5mwm .rough-node .label text,#mermaid-svg-H6FqUY7dgiws5mwm .node .label text,#mermaid-svg-H6FqUY7dgiws5mwm .image-shape .label,#mermaid-svg-H6FqUY7dgiws5mwm .icon-shape .label{text-anchor:middle;}#mermaid-svg-H6FqUY7dgiws5mwm .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-H6FqUY7dgiws5mwm .rough-node .label,#mermaid-svg-H6FqUY7dgiws5mwm .node .label,#mermaid-svg-H6FqUY7dgiws5mwm .image-shape .label,#mermaid-svg-H6FqUY7dgiws5mwm .icon-shape .label{text-align:center;}#mermaid-svg-H6FqUY7dgiws5mwm .node.clickable{cursor:pointer;}#mermaid-svg-H6FqUY7dgiws5mwm .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-H6FqUY7dgiws5mwm .arrowheadPath{fill:#333333;}#mermaid-svg-H6FqUY7dgiws5mwm .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-H6FqUY7dgiws5mwm .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-H6FqUY7dgiws5mwm .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-H6FqUY7dgiws5mwm .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-H6FqUY7dgiws5mwm .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-H6FqUY7dgiws5mwm .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-H6FqUY7dgiws5mwm .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-H6FqUY7dgiws5mwm .cluster text{fill:#333;}#mermaid-svg-H6FqUY7dgiws5mwm .cluster span{color:#333;}#mermaid-svg-H6FqUY7dgiws5mwm div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-H6FqUY7dgiws5mwm .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-H6FqUY7dgiws5mwm rect.text{fill:none;stroke-width:0;}#mermaid-svg-H6FqUY7dgiws5mwm .icon-shape,#mermaid-svg-H6FqUY7dgiws5mwm .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-H6FqUY7dgiws5mwm .icon-shape p,#mermaid-svg-H6FqUY7dgiws5mwm .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-H6FqUY7dgiws5mwm .icon-shape .label rect,#mermaid-svg-H6FqUY7dgiws5mwm .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-H6FqUY7dgiws5mwm .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-H6FqUY7dgiws5mwm .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-H6FqUY7dgiws5mwm :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 残差近路:原样抄一条
残差
输入 x 50, 64
LayerNorm(先换标准分)
① 多头自注意力(开会)

nhead=4 个小组同时讨论
x + Attn(LN(x))
LayerNorm
② 前馈网络(回工位消化)

Linear 64→128 → ReLU → Linear 128→64
送入下一层 × 共 2 层

第二轮会在第一轮结论上继续
出口 LayerNorm

⊕ 是残差相加点:交流结果只是「补充意见」加回原来的自己,就算这层学砸了,信息也能原封不动穿过。

5. 分类头:班长投票

取出第 0 号位置的班长(tokens[:, 0]),经 LayerNorm 后用一个 Linear(64→10) 加权求和,压成 10 个分数(logit),谁分最高就预测成哪个数字。

四、完整数据流

步骤 操作 输出形状
1 输入灰度图 [B, 1, 28, 28]
2 两层卷积 + BN + ReLU + 池化 [B, 64, 7, 7]
3 1×1 卷积混合通道到 d_model [B, 64, 7, 7]
4 展平成 49 个 token + 拼接 CLS [B, 50, 64]
5 2 层 Transformer Encoder [B, 50, 64]
6 取 CLS 输出 → 全连接分 10 类 [B, 10]

CNN vs Transformer 一句话总结:

CNN Transformer
看什么 局部笔画 全局关系
视角 放大镜扫过每个小窗口 每块拼图看所有其他块
输出 7×7×64 特征图 50 个 64 维向量
模型里的角色 提取特征 让特征互相沟通、汇总

五、目录结构

复制代码
mnist-cnn-transformer/
├── config.py      # 全部超参数(batch、lr、维度...)
├── model.py       # CNNTransformer 网络定义
├── train.py       # 训练、验证、保存权重
├── test.py        # 加载权重,测试集评估
├── web/           # 本地手写识别网页(Flask)
├── visualize/     # Torchview 模型结构图
├── tensorboard/   # TensorBoard 训练曲线
├── flow/          # 单张图的数据流故事板
└── test_flow/     # 整场测试总览 + 混淆矩阵

五个扩展目录各自独立、不改动核心代码,需要哪种可视化就装哪种依赖。

六、快速上手

bash 复制代码
# 1. 创建环境(Python 3.10,仅 CPU)
conda env create -f environment.yml
conda activate mnist-cnn-transformer

# 2. 训练(首次运行自动下载 MNIST)
python train.py --epochs 8 --batch-size 64

# 3. 测试
python test.py

训练中断可续训:

bash 复制代码
python train.py --resume --epochs 8   # --epochs 表示"练到第 8 轮",不是"再练 8 轮"

七、可视化玩法

① 本地手写识别(网页)

bash 复制代码
pip install -r web/requirements.txt
python web/app.py     # 打开 http://127.0.0.1:5000

在画板上手写数字,松手即识别,并实时展示:

  • 预测结果 + 10 类概率条;
  • CNN 各层特征图(不同卷积通道);
  • Transformer 每层 7×7 空间向量箭头(PCA 投影)。

② 模型结构图

bash 复制代码
python visualize/export_graph.py    # 需安装 Graphviz

③ TensorBoard 曲线

bash 复制代码
python tensorboard/train_tb.py --epochs 8
tensorboard --logdir tensorboard/runs

④ 数据流故事板

bash 复制代码
python flow/export_flow.py --digit 7        # 单张图从输入到预测
python test_flow/export_test.py             # 整场测试总览 + 对/错样例

八、项目地址

🔗 源码:https://gitee.com/wufengsheng/mnist-cnn-transformer

欢迎 Star、提 Issue,一起学习深度学习!

相关推荐
澄旭13 分钟前
不同 AI Agent 共用同一套 Skills
人工智能
阿牛哥_GX15 分钟前
接入AI大模型:让机器人"智能"起来
人工智能
元启数宇26 分钟前
幕墙AI设计算法逻辑:元启数宇如何智能分格
人工智能·算法
故七月38 分钟前
以合规化技术体系筑牢GEO产业发展根基 万域智瞰引领AI流量服务高质量发展
大数据·人工智能
开发小程序的之朴43 分钟前
从神经网络到文件加密:一次关于 SIREN、ARX 与密码安全性的实验探索
人工智能·深度学习·神经网络
哈基咩1 小时前
Function Calling 与 Code Mode:AI Agent 工具调用的原理、对比与选型指南
网络·人工智能
就是一顿骚操作1 小时前
ViT:把图像切成词之后,Transformer 如何进入计算机视觉
人工智能·深度学习·计算机视觉·transformer·论文解读
千里码aicood1 小时前
PyQt基于卷积神经网络的智慧校园的设计与实现
人工智能·cnn·pyqt
MartinYeung51 小时前
[论文学习]MPIB:医疗提示注入基准——针对LLM临床安全性的系统性评估
人工智能·python·学习