发布时间:2026/8/30 12:09:43
从零实现Vision Transformer:ViT原理与PyTorch代码详解 Transformer 目前是深度学习中最具影响力的序列建模架构之一它的核心思想来自 2017 年的论文 Attention Is All You Need用自注意力机制直接计算序列中任意两个位置之间的相关性不再依赖循环或者卷积逐步传递信息。Vision TransformerViT把这个架构搬到了图像任务上思路很直接先把图像切成一堆固定大小的 patch再把每个 patch 拉平并投影成一个向量也就是 Patch Embedding然后把这些向量当作 token 序列输入标准 Transformer Encoder最后用分类头输出结果。真正理解 ViT不能只停留在会调用现成库的层面还要能解释清楚一个输入图像从进入网络到输出 logits 的过程中每个张量的形状发生了什么变化每个模块为什么这么设计。这篇文章会从 Transformer 的核心机制讲起然后逐步实现 Patch Embedding、多头自注意力、Transformer Encoder Block 和完整 ViT 模型最后用 MNIST 分类任务把整条链路跑通并给出训练、验证、排错和工程化建议。1. Transformer 到底在解决什么问题1.1 RNN 和 CNN 的局限在 Transformer 出现之前序列建模最常用的是 RNN 和它的变体 LSTM、GRU。RNN 把序列元素按时间顺序逐个读入维护一个隐状态当前时刻的输出依赖上一时刻的隐状态。这种方式有两个明显问题第一长距离信息要经过很多次隐状态传播才能到达目标位置中间容易出现梯度消失或信息衰减第二时间步之间是串行计算无法像矩阵运算一样充分并行训练效率受到限制。CNN 在视觉任务上是主力。它通过局部感受野和卷积核共享权重天然带有平移不变性和局部性的归纳偏置。但是要建模整张图像的全局关系CNN 需要堆叠大量卷积层让感受野逐层扩大这导致建模远程依赖的成本比较高。而且远距离的两个像素是否相关并不一定与它们的空间距离成正比。CNN 的局部优先假设在 ImageNet 这类大数据上仍然有效但模型结构本身给全局建模带来了额外负担。Transformer 和它们最本质的区别是它不假设信息必须沿着时间顺序或空间邻接关系传递。对于输入序列中的任意两个元素Transformer 都计算一个注意力权重权重越高表示目标元素在更新自身表示时越依赖那个元素。这样全局关系在每一层都是显式建模的而且所有位置可以并行计算。1.2 自注意力机制的工作原理自注意力解决的问题是给定一个序列让序列中的每个元素都能根据全序列其他元素的信息来更新自己。可以把输入序列理解为 N 个 token每个 token 是一个向量例如形状是[N, D]。为了让每个 token 与其他 token 交互模型为每个 token 生成三个向量Query、Key、Value。Query 表示“我想找什么信息”Key 表示“我能提供什么信息”Value 表示“我实际携带的信息”。某个 token 的 Query 与所有 token 的 Key 做点积得到一个相关性分数分数经过 softmax 变成权重再用权重对所有 Value 做加权求和得到该 token 更新后的输出。写成公式就是Attention(Q, K, V) softmax(Q K^T / sqrt(d_k)) V其中d_k是每个注意力头的维度。除以sqrt(d_k)是为了避免Q K^T的数值随维度增大而变得太大导致 softmax 进入饱和区梯度变得非常小。多头注意力的做法是把 D 维的 Query、Key、Value 分别拆成 H 组每组维度为 D/H各自独立做注意力计算最后拼接起来再过一次线性投影。这样可以让模型在多个子空间里分别关注不同类型的依赖关系有的头可能关注相邻 patch有的头可能关注全局轮廓。1.3 为什么图像也要用 TransformerViT 的思路不是用卷积提取特征后再接 Transformer而是直接把图像变成 token 序列。它的动机是如果数据量足够大模型对“局部性”的人工预设并不是必需的。Transformer 可以通过注意力自己学到哪些 patch 需要相互关注。不过这里有一个重要的边界在中小规模数据集上直接从头训练 ViT效果通常不如同规模的 CNN。原因在于 CNN 的归纳偏置在数据不足时是一种保护而 Transformer 更依赖大规模数据来学习结构。所以个人项目和工业落地时常见做法是使用在大规模数据上预训练好的 ViT 权重做迁移学习而不是从随机初始化开始训练。2. ViT 的整体架构图像怎么变成 token 序列2.1 ViT 处理图像的完整流水线假设输入一张 RGB 图像形状是[B, 3, 224, 224]其中 B 是 batch size。ViT 的流水线如下图像切 patch把 224x224 的图像按 16x16 的 patch 划分得到 14x14196 个 patch。Patch Embedding把每个 patch 投影成一个 D 维向量得到[B, 196, D]的 token 序列。拼接 CLS token在序列开头添加一个特殊的可学习 token得到[B, 197, D]。加位置编码加上一个[B, 197, D]的 position embedding保持长度不变。过 L 层 Transformer Encoder每层都是多头自注意力 MLP LayerNorm 残差。取 CLS token 的输出[B, D]。分类头线性层映射到类别数得到[B, num_classes]。如果任务是目标检测或分割第 6 步会替换成其他解码器或特征图输出结构但前面第 1 到第 5 步基本相同。2.2 Patch Embedding把图像切块并投影成向量图像是一个高维数组不能直接当成一维 token 序列输入 Transformer。Patch Embedding 做的事情就是把“二维图像上的一个小块”编码成“一个向量”。具体来说一个大小为patch_size x patch_size、通道数为 C 的 patch拉平后长度是patch_size^2 * C。用一个线性层可以把这个向量投影到 embed_dim 维。实际操作中切块加投影可以用一个二维卷积一步完成卷积核大小和步长都等于 patch_size输入通道数等于 C输出通道数等于 embed_dim。卷积输出的特征图上每个位置对应原图的一个 patch且每个通道上的数值就是这个 patch 在该维度上的投影结果。例如输入[B, 3, 224, 224]patch_size16embed_dim768卷积输出就是[B, 768, 14, 14]。把它展平成[B, 768, 196]再转置成[B, 196, 768]就得到了 token 序列。这里每个 token 对应原图 16x16 的一个 patch。2.3 Position Embedding、CLS token 与分类头注意力机制对输入的顺序不敏感。如果把 patch 序列任意打乱自注意力计算结果只是行顺序变化数值不会改变。但图像的空间顺序是有意义的所以必须把位置信息注入。ViT 使用可学习的位置编码初始化后随训练一起更新。它的形状是[1, num_patches 1, embed_dim]因为前面还要拼接一个 CLS token。也有人会用正弦位置编码但 ViT 论文和工程实现大多是直接学习。CLS token 是拼接在 patch 序列开头的一个可学习向量。它没有对应的输入 patch只是作为全图信息的汇聚点。经过多层 Encoder 后CLS 位置的输出向量通过自注意力汇总了所有 patch 的信息因此可以用它来分类。关于 CLS token 和全局平均池化ViT 里两种做法都有人用。CLS token 的好处是输出与序列长度解耦而且和预训练结构一致平均池化则更显式地利用所有 patch 信息。在迁移学习中尽量保持与预训练一致的用法。3. 环境准备与项目结构3.1 环境依赖演示代码基于 Python 3.9 和 PyTorch 2.x 编写主要依赖如下依赖版本建议作用Python3.8运行环境PyTorch2.0张量计算与自动求导torchvision0.15数据集与图像变换numpy1.24间接依赖安装命令conda create -n vit-demo python3.9 -y conda activate vit-demo pip install torch torchvision如果没有 GPU安装 CPU 版即可。本文的最小演示模型很小CPU 上几个 epoch 也能跑完。如果需要 GPU 训练安装对应 CUDA 版本的 PyTorch。3.2 项目结构vit-from-scratch/ ├── vit.py # 模型定义PatchEmbed、Attention、Encoder、ViT ├── train.py # 数据加载与训练验证 └── data/ # 数据目录vit.py是核心所有模型组件都放在里面。train.py负责把 MNIST 数据加载进来创建模型执行训练和验证。3.3 演示数据集选择为了在普通机器上快速验证本文使用 MNIST 手写数字分类。MNIST 是 28x28 的单通道灰度图共 10 类。我们会把它 Resize 到 32x32patch_size4这样会得到 64 个 patch训练速度很快。如果读者想换 CIFAR-10只需要把in_channels改成 3并修改 Resize 和归一化参数想换 ImageNet 风格的图片需要调整img_size和patch_size保证两者可以整除。4. 代码手撕从零实现 Vision Transformer这一节是实现重点。这里的“手撕”不是把开源代码抄一遍而是用一个最小但完整的实现展示 ViT 的每个模块。模型定义全部放在vit.py中不使用 torchvision 或 timm 里现成的 ViT。4.1 PatchEmbed图像切块与线性投影import torch import torch.nn as nn class PatchEmbed(nn.Module): 将图像切成 patch并线性投影为 token 向量。 输入: [B, C, H, W] 输出: [B, num_patches, embed_dim] def __init__(self, img_size32, patch_size4, in_channels1, embed_dim128): super().__init__() if img_size % patch_size ! 0: raise ValueError( fimg_size{img_size} 必须能被 patch_size{patch_size} 整除 ) self.img_size img_size self.patch_size patch_size self.num_patches (img_size // patch_size) ** 2 # 卷积核大小和步长都等于 patch_size # 等价于先把每个 patch 拉平再过一个线性层 self.proj nn.Conv2d( in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size, ) def forward(self, x): B, C, H, W x.shape # 检查输入尺寸避免隐含错误在后面的位置编码阶段才暴露 if H ! self.img_size or W ! self.img_size: raise ValueError( f输入尺寸应为 {self.img_size}x{self.img_size}实际为 {H}x{W} ) # [B, embed_dim, H/patch_size, W/patch_size] x self.proj(x) # [B, embed_dim, num_patches] x x.flatten(2) # [B, num_patches, embed_dim] x x.transpose(1, 2) return x这个模块的核心是nn.Conv2d(in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size)。卷积核在图像上以 patch_size 为步长滑动每个位置对应一个 patch输出通道数就是 embedding 维度。这样做的好处是切块和线性投影在同一个操作里完成计算效率高也不需要手动写torch.nn.functional.unfold。4.2 MultiHeadSelfAttention多头自注意力import math import torch import torch.nn as nn import torch.nn.functional as F class MultiHeadSelfAttention(nn.Module): 多头自注意力。 输入: [B, N, D] 输出: [B, N, D] def __init__(self, embed_dim128, num_heads4, attn_dropout0.0): super().__init__

相关新闻

2026/8/30 12:09:43

OpenClaw走向LTS:智能体部署、Skill开发与本地模型实践

OpenClaw: On the Road to LTS——智能体工具开始认真谈“长期支持”了如果你最近在关注开源 AI Agent 项目,大概会注意到一个高频词:LTS。过去我们说 LTS,更多是在说 Ubuntu、Spring Boot、Node.js 这些“基础软件”;而现在&…

2026/8/30 12:04:42

从美团2013笔试卷看大厂研发岗基础能力考察

老有人问我,十年前的美团笔试到底考什么,和现在的互联网公司笔试有什么区别。我也翻出过不少老题,其中“美团2013湖南研发工程师笔试卷”这份流传很广的卷子,确实是观察国内互联网研发岗早期面试风格的一个不错样本。那时候移动互…

2026/8/30 12:19:44

UltraPIPS:基础模型如何提升超声图像感知能力

1. 背景与核心概念 医学超声成像一直是一个“入门容易、精通难”的领域。相比 CT 和 MRI,B 型超声(B-mode ultrasound)没有电离辐射、检查成本低、实时性好,因此在腹部、心血管、甲状腺、乳腺以及产科检查中占有不可替代的地位。但…

2026/8/30 12:19:44

STM32与边缘AI:从传感器到网关的工业IoT落地指南

拿到《23STM32峰会资料》意法半导体IoT助力工业智能化这份材料,我的第一反应是:ST这次没有把工业物联网讲成PPT上的“大趋势”,而是给了一条从传感器、MCU、MPU、无线连接一直到边缘AI工具链的落地链路。过去大家聊STM32,注意力基…

2026/8/30 12:19:44

Java面试八股文精简版:高频考点与加分答法

我做了这么多年Java面试辅导,也当过不少次面试官,最大的感受就是:很多候选人不是不会,而是不知道面试官到底想听什么。网上八股文一抓一大把,但质量参差不齐,有的过于啰嗦,有的又漏掉了真正的考…

2026/8/30 12:19:44

book-to-skill:把书籍蒸馏成AI Skill,用Pygame开发小游戏

book-to-skill 是一种把书籍资料转化为 AI Skill 的工作方式:先拆解一本书里的知识结构,提取出可执行的规则、流程和模板,再把这些内容打包成让 AI 编程工具可以直接调用的 Skill。它的价值在游戏开发场景里特别明显,因为游戏开发…

2026/8/30 12:19:44

TouchGFX从4.13到4.18升级实践:嵌入式UI版本迁移全指南

半年前交付的HMI项目最近又要加新功能,打开旧工程时发现TouchGFX Designer已经提示“Project created with an older version of TouchGFX Designer”。一开始我没当回事,以为点个继续就行,结果客户发来的新UI资源在旧版本里压根打不开。Touc…

2026/8/30 12:14:44

Zsh历史记录增强实战:智能搜索、去重与多终端同步

这次我们来看一个很有意思的开发者项目: Show HN: Smarter Shell History for Zsh 。简单说,它不是又一个“用 AI 生成命令”的玩具,而是把 Zsh 自带的历史记录能力重新做了一层增强,让终端里翻历史、找上次跑过的命令、批量复用…

2026/8/30 0:03:35

vSound小提琴数字处理器实操指南:从接线到演出的完整配置

电小提琴或者原声小提琴插电演出,第一个绕不开的坎就是声音难听。原声琴的共鸣和空气感一旦进了拾音器,出来的往往是一坨干瘪、发尖、带着奇怪塑料味的信号。我当初第一次把琴接上乐队调音台,直接被主唱吐槽"你这声音像在锯钢丝"。…

2026/8/30 0:03:35

传感器接口IC如何攻克生物化学传感的微弱信号难题?

1. 从电极到比特流:为什么生物化学传感必须依赖专用接口IC 做生物化学传感的人都有过类似的经历:明明传感器本身性能很好,信号输出却一塌糊涂——噪声大、漂移明显、重复性差,怎么调都达不到预期。很多时候问题并不在传感器&#…

2026/8/30 0:03:35

STM32F411CEU6多通道ADC采集:扫描模式+DMA实现详解

1. 多通道 ADC 的用武之地把“Multichannel ADC”和“STM32F411CEU6”这两个关键字放在一起,其实就是嵌入式开发里最常遇到的一类需求:用一块不算贵的 MCU,同时采集多路模拟信号。STM32F411CEU6 是 48 引脚的 Cortex-M4F 主控,主频…

2026/8/30 0:03:35

vSound小提琴数字处理器实操指南:从接线到演出的完整配置

电小提琴或者原声小提琴插电演出,第一个绕不开的坎就是声音难听。原声琴的共鸣和空气感一旦进了拾音器,出来的往往是一坨干瘪、发尖、带着奇怪塑料味的信号。我当初第一次把琴接上乐队调音台,直接被主唱吐槽"你这声音像在锯钢丝"。…

2026/8/30 0:03:35

传感器接口IC如何攻克生物化学传感的微弱信号难题?

1. 从电极到比特流:为什么生物化学传感必须依赖专用接口IC 做生物化学传感的人都有过类似的经历:明明传感器本身性能很好,信号输出却一塌糊涂——噪声大、漂移明显、重复性差,怎么调都达不到预期。很多时候问题并不在传感器&#…

2026/8/30 0:03:35

STM32F411CEU6多通道ADC采集:扫描模式+DMA实现详解

1. 多通道 ADC 的用武之地把“Multichannel ADC”和“STM32F411CEU6”这两个关键字放在一起,其实就是嵌入式开发里最常遇到的一类需求:用一块不算贵的 MCU,同时采集多路模拟信号。STM32F411CEU6 是 48 引脚的 Cortex-M4F 主控,主频…

2026/8/28 16:16:48

实测才敢推 AI论文网站 2026最新测评与推荐

2026年真正好用的AI论文网站,核心看生成的论文质量、低AI味、格式正确、学术适配四大指标。综合实测,千笔AI、ThouPen、豆包、DeepSeek、Grammarly 是当前最值得推荐的梯队,覆盖从免费到付费、从中文到英文、从文科到理工的全场景需求。一、综…

2026/8/28 16:16:50

2026必备!AI论文网站测评:最新推荐与深度对比

2026年真正好用的AI论文网站,核心看生成的论文质量、低AI味、格式正确、学术适配四大指标。综合实测,千笔AI、ThouPen、豆包、DeepSeek、Grammarly 是当前最值得推荐的梯队,覆盖从免费到付费、从中文到英文、从文科到理工的全场景需求。 一、…

2026/8/28 11:06:45

摆脱论文困扰!盘点2026年全网爆红的的AI论文写作工具

一天写完毕业论文在2026年已不再是天方夜谭。2026年最炸裂、实测能大幅提速的AI论文写作工具,覆盖选题构思、文献整理、内容生成、格式排版等核心场景,真正帮你高效搞定论文难题。 一、全流程王者:一站式搞定论文全链路(一天定稿首…