Vision Transformer编码流程及代码详解
前言
传统CNN依靠卷积核局部滑动提取图像特征,依赖归纳偏置(局部性、平移不变性);而Vision Transformer(ViT)完全基于自注意力机制,将图像拆分为序列Patch,借用NLP Transformer架构完成全局特征建模。
本文以工程最常用的ViT-B/16为例,完整拆解从输入图像[3,224,224]到最终一维图像特征[1,768]的全流程维度变换、核心模块原理,并附带可运行PyTorch完整实现代码,逐行注释方便调试。
一、ViT-B/16 参数含义
B表示Base,代表模型基础尺寸,是轻量化常用版本,编码器堆叠12层;
- L(Large):大模型,编码器24层
- H(Huge):超大模型,编码器32层
16代表Patch分块尺寸:将224×224原图切分为16×16像素的小块;
常见Patch尺寸:16、32、14(ViT-L/14多用于高精度图像任务)。
二、前置基础:核心模块维度运算基础
理解ViT的关键是全程跟踪张量维度变化,先回顾矩阵乘法规则:
向量a = [1,768]× 权重矩阵W = [768,512]= 输出向量c = [1,512]
内维度必须相等,输出维度由向量第一维、矩阵第二维决定。
下面逐个介绍ViT编码器全部基础组件:
1. Linear 线性层
公式:y=Wx+by = Wx + by=Wx+b
对输入特征做线性投影,无非线性,多用于Patch嵌入、多头注意力映射、分类头映射,是维度升降维的核心层。
2. GELU 激活函数
高斯误差线性单元,替代ReLU,平滑非线性激活。
解决ReLU梯度硬截断问题,ViT前馈网络统一使用GELU提升特征拟合能力,捕获图像复杂纹理、语义特征。
3. Dropout 随机失活
正则化手段,训练阶段随机置零部分神经元输出,推理阶段恢复完整权重。
迫使模型不依赖局部少数神经元,降低过拟合,ViT在嵌入层、注意力输出、前馈层后均会添加Dropout。
4. Layer Normalization 层归一化
对单一样本自身特征维度做归一化,区别于CNN常用的BatchNorm(批次维度归一)。
Transformer系列标配,稳定每层输入分布,大幅加速深层模型收敛,每层注意力、前馈网络前都会先做LN。
5. Self-Attention 自注意力机制
核心:计算序列内每个Token与所有Token的相关性权重。
输入序列中每个Patch Token互相计算相似度,建模图像全局依赖(远距离像素关联,CNN很难做到)。
6. Multi-Head Attention 多头自注意力
将特征通道均分N个注意力头,每个头独立计算自注意力,最后拼接所有头输出再线性融合。
多个头并行捕捉不同维度、不同尺度的空间关联(边缘、色块、全局轮廓),单头注意力表达能力不足。
7. FFN 前馈神经网络
两层全连接+GELU激活:升维映射→激活→降维映射。
独立作用于序列每一个Token,对注意力输出特征做非线性特征变换,增强模型表征能力。
8. Residual Connection 残差连接
Output=Input+SubLayer(Input)Output = Input + SubLayer(Input)Output=Input+SubLayer(Input)
每层注意力、FFN外层包裹残差相加,深层堆叠时避免梯度消失,保证梯度跨层回流,是12/24层深层ViT训练的基础。
9. Positional Encoding 位置编码
自注意力本身不感知序列顺序,图像Patch打散后丢失空间位置信息。
通过可学习位置编码(ViT原生方案)或正余弦编码,生成和Patch Embedding同维度位置向量,逐元素相加嵌入序列,还原图像二维空间信息。
三、ViT完整编码工作流程(维度全程跟踪)
输入:单张RGB图像张量[C=3, H=224, W=224],batch_size=1,输入形状[1,3,224,224]
步骤1:图像切分Patch块
按照Patch_size=16切割原图:
横向分块:224/16=14224 / 16 = 14224/16=14,纵向同理14块,总Patch数量14×14=19614×14=19614×14=196
每个Patch像素尺寸:[3,16,16]
全局图像拆分后得到196个独立图像小块。
步骤2:Patch Embedding 图像分块嵌入
- 单个Patch
[3,16,16]展平一维:3×16×16=7683×16×16=7683×16×16=768,单Patch展平向量维度[768] - 全部196个Patch堆叠,得到原始Patch序列:
[196, 768] - 批量维度扩展:batch=1,张量形状
[1, 196, 768]
核心逻辑:使用Conv2d卷积等价实现Patch切分+线性投影(工程上速度更快)
卷积核=16,步长=16,输出通道768,卷积输出直接reshape为序列Token。
步骤3:拼接Class Token分类向量
ViT新增一个专属分类Token(Class Token),形状[1,1,768],拼接在Patch序列最前端:
原序列[1,196,768]+ Class Token → 新序列[1, 197, 768]
模型最终全局图像特征从该Class Token提取,对应文末输出[1,768]特征。
步骤4:叠加可学习位置编码
创建与序列等长的位置编码参数[197,768](覆盖196个Patch + 1个Class Token),逐元素相加到嵌入序列,注入空间位置信息。
叠加后张量尺寸不变:[1, 197, 768],随后经过Dropout做正则。
步骤5:堆叠12层Transformer Encoder(ViT-B核心编码层)
每层Encoder结构固定:LN层 → 多头自注意力 + 残差连接 → LN层 → FFN前馈网络 + 残差连接
逐层迭代计算全局注意力特征,每层输入输出维度均保持[1,197,768]不变。
单层Encoder数据流:
- 输入
x = [1,197,768] - 层归一化LN1 → 多头注意力MHA → x_attn = x + MHA(LN1(x)) 残差相加
- 对x_attn做层归一化LN2 → FFN前馈网络 → x_out = x_attn + FFN(LN2(x_attn)) 残差相加
- x_out作为下一层Encoder输入
12层循环结束,最终输出编码后完整序列[1, 197, 768]
步骤6:提取全局图像特征(目标输出[1,768])
197个Token中,第0位为Class Token,代表整张图像聚合全局语义特征:
切片取出x[:, 0, :],形状[1, 768],即文章开头所说图像全局特征向量。
后续分类任务可再接Linear层映射至类别数量,检测/分割任务则取用全部Patch Token[:,1:,:]。
四、ViT-B/16 完整PyTorch实现代码
importtorchimporttorch.nnasnnimporttorch.nn.functionalasF# 超参数配置 ViT-B/16BATCH_SIZE=1IMG_CHANNEL=3IMG_SIZE=224PATCH_SIZE=16EMBED_DIM=768# Base模型特征维度NUM_HEADS=12# 多头注意力头数NUM_LAYERS=12# Encoder层数MLP_HIDDEN=3072# FFN隐藏层维度DROPOUT_RATE=0.1# 1. 单层FFN前馈网络classFeedForward(nn.Module):def__init__(self):super().__init__()self.net=nn.Sequential(nn.Linear(EMBED_DIM,MLP_HIDDEN),nn.GELU(),nn.Dropout(DROPOUT_RATE),nn.Linear(MLP_HIDDEN,EMBED_DIM),nn.Dropout(DROPOUT_RATE))defforward(self,x):returnself.net(x)# 2. 单层Transformer EncoderclassTransformerEncoderLayer(nn.Module):def__init__(self):super().__init__()self.norm1=nn.LayerNorm(EMBED_DIM)self.attn=nn.MultiheadAttention(EMBED_DIM,NUM_HEADS,dropout=DROPOUT_RATE,batch_first=True)self.norm2=nn.LayerNorm(EMBED_DIM)self.ffn=FeedForward()defforward(self,x):# 多头注意力 + 残差attn_out,_=self.attn(query=self.norm1(x),key=self.norm1(x),value=self.norm1(x))x=x+attn_out# FFN + 残差ffn_out=self.ffn(self.norm2(x))x=x+ffn_outreturnx# 3. 完整ViT-B/16 编码器classViT_B16_Encoder(nn.Module):def__init__(self):super().__init__()num_patches=(IMG_SIZE//PATCH_SIZE)**2# 14*14=196# Patch Embedding:卷积替代分块+线性投影self.patch_embed=nn.Conv2d(IMG_CHANNEL,EMBED_DIM,kernel_size=PATCH_SIZE,stride=PATCH_SIZE)# 可学习Class Tokenself.cls_token=nn.Parameter(torch.randn(1,1,EMBED_DIM))# 可学习位置编码:196patch + 1cls_tokenself.pos_embed=nn.Parameter(torch.randn(1,num_patches+1,EMBED_DIM))self.pos_drop=nn.Dropout(DROPOUT_RATE)# 堆叠12层Encoderself.encoder_layers=nn.Sequential(*[TransformerEncoderLayer()for_inrange(NUM_LAYERS)])self.norm_final=nn.LayerNorm(EMBED_DIM)defforward(self,img):# img输入 shape [B,3,224,224]B=img.shape[0]# Step1 Patch Embedding [B,768,14,14] -> [B,196,768]patch_feat=self.patch_embed(img)patch_feat=patch_feat.flatten(2).transpose(1,2)# Step2 拼接Class Tokencls_tokens=self.cls_token.expand(B,-1,-1)# [B,1,768]x=torch.cat([cls_tokens,patch_feat],dim=1)# [B,197,768]# Step3 叠加位置编码+Dropoutx=x+self.pos_embed x=self.pos_drop(x)# Step4 12层Transformer编码x=self.encoder_layers(x)x=self.norm_final(x)# Step5 提取全局图像特征 cls_token [B,768]global_img_feat=x[:,0,:]returnglobal_img_feat# 测试流程if__name__=="__main__":# 模拟输入图片 [1,3,224,224]test_img=torch.randn(BATCH_SIZE,IMG_CHANNEL,IMG_SIZE,IMG_SIZE)model=ViT_B16_Encoder()feat_out=model(test_img)print("输入图像尺寸:",test_img.shape)print("输出全局图像特征尺寸:",feat_out.shape)# 输出结果:torch.Size([1, 768]),和文中结论完全对应