尧图网站建设 尧图网络
  • 首页
  • 关于我们
  • 服务项目
  • 案例展示
  • 建站流程
  • 资讯中心
  • 联系我们
首页/资讯中心/详情

Vision Transformer编码流程及代码详解

Vision Transformer编码流程及代码详解
📅 发布时间:2026/7/22 12:39:59

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 图像分块嵌入

  1. 单个Patch[3,16,16]展平一维:3×16×16=7683×16×16=7683×16×16=768,单Patch展平向量维度[768]
  2. 全部196个Patch堆叠,得到原始Patch序列:[196, 768]
  3. 批量维度扩展: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数据流:

  1. 输入x = [1,197,768]
  2. 层归一化LN1 → 多头注意力MHA → x_attn = x + MHA(LN1(x)) 残差相加
  3. 对x_attn做层归一化LN2 → FFN前馈网络 → x_out = x_attn + FFN(LN2(x_attn)) 残差相加
  4. 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]),和文中结论完全对应

相关新闻

  • TM4C129 I2C中断机制详解:从寄存器配置到实战优化
  • 手机AI修图技术解析:从原理到实战应用
  • 标识工程出片品质哪家高?2026年十大出片品牌深度测评,所见即所得不踩雷 - 工业品牌热点

最新新闻

  • 青岛做水产养殖加工的客户,AI获客系统怎么帮您应对客户犹豫? - 红枫叶GEO优化公司
  • 权威核验|2026年7 月江诗丹顿售后服务中心实地考察网点地址+电话全更新 - 江诗丹顿中国服务中心
  • 万息投标,标书查重与审查工具,让每一份标书都经得起检验
  • 江诗丹顿售后网点核验报告|最新维修地址及电话权威收录(2026年7月最新) - 江诗丹顿中国服务中心
  • Qwen3模型Lora微调实战:基于LLaMA-Factory的高效方案
  • 智能测绘装备赛道升温:扫描全站仪行业市场格局、痛点机遇与发展前瞻

日新闻

  • AI云原生实战05-金融AI上云最难的不是技术,是“不出事“——TCE银行风控架构拆解
  • 2026年GEOSEO优化公司选型深度测评:五大硬核标准严选,这六家重塑搜索增长新格局 - 品牌前沿专家
  • **核验!2026年7月卡地亚香港**售后网点地址及服务电话公告 - 卡地亚服务中心

周新闻

  • SaaS软件行业GEO实践:AI搜索时代的品牌可见性与获客新路径
  • 什么是PCTFE?医药高端包装的“防潮王牌“材料
  • 【JVM调优实战】16-可视化利器-JConsole-VisualVM-JMC

月新闻

  • 2026年6月公司网站搭建最新热门渠道测评:四大低成本/零代码平台对比+避坑
  • 【Linux】Linux arm 编译QT程序,出现expected “}“报错
  • 【MATLAB例程】四基站二维AOA定位与距离辅助增强对比仿真。基于角度观测和测距修正的固定目标平面定位精度分析

关于尧图

  • 公司简介
  • 团队介绍
  • 企业文化
  • 荣誉资质

服务项目

  • 定制开发
  • 电商建站
  • UI 设计
  • 运维服务

快速链接

  • 案例展示
  • 建站流程
  • 常见问题
  • 资讯中心

联系方式

  • 📍北京市朝阳区互联网产业园 A 座 10 层
  • 📞400-888-8888
  • ✉️contact@rkmt.cn
  • 🕐周一至周日 9:00-21:00

© 2024 北京尧图网络科技有限公司 版权所有 | 京 ICP 备 XXXXXXXX 号