学习 VLA 第5天:VIT算法原理以及复现
1. 引言

本文是学习 VLA(Vision-Language-Action,视觉-语言-动作)模型的第 5 天内容。前几天的学习中,已经了解了 VLA 模型的整体架构、多模态大模型的基础知识,以及视觉编码器在 VLA 中的重要作用。本文将深入剖析 VLA 模型中视觉编码器的核心组件——ViT(Vision Transformer,视觉 Transformer)的算法原理,并动手实现一个简化版的 ViT 模型。
2. ViT 的核心思想
2.1 从 NLP 到 CV 的迁移
ViT 的核心思想非常简洁:将图像切分成固定大小的 Patch(图像块),然后把每个 Patch 展平并线性投影为向量,最后将这些向量作为 Token 输入标准的 Transformer 编码器。
这与 NLP 中把句子切分成词元(Token)的思路完全一致。Transformer 本身并不关心输入是文本还是图像,它只处理一组向量序列。因此,只要能将图像转化为一组向量序列,就能直接复用 Transformer 架构。
2.2 为什么需要 ViT

在 ViT 之前,视觉任务几乎被 CNN(卷积神经网络)统治。CNN 通过卷积核在图像上滑动来提取局部特征,具有平移等变性和局部感受野等归纳偏置。然而,CNN 的局部感受野限制了它对全局信息的建模能力,需要堆叠大量卷积层才能扩大感受野。
ViT 则完全抛弃了卷积操作,通过自注意力机制直接建模图像中任意两个位置之间的关系,天然具备全局建模能力。在数据量足够大的情况下,ViT 的性能可以超越 CNN。
下面用一张流程图概括 ViT 的核心思想:
3. ViT 的模型结构
ViT 的整体结构可以分为以下几个部分:
- Patch Embedding(图像块嵌入):将图像切分为 Patch 并投影为向量。
- 位置编码(Position Embedding):为每个 Patch 添加位置信息。
- [CLS] Token:用于聚合全局信息的特殊 Token。
- Transformer Encoder(Transformer 编码器):多层自注意力 + 前馈网络。
- 分类头(Classification Head):用于输出分类结果。
下面逐一讲解每个部分。
3.1 Patch Embedding
假设输入图像大小为 H×W×CH \times W \times CH×W×C(高 × 宽 × 通道数),将其切分为大小为 P×PP \times PP×P 的 Patch。那么一共可以得到 N=H×WP2N = \frac{H \times W}{P^2}N=P2H×W 个 Patch。
每个 Patch 的形状为 P×P×CP \times P \times CP×P×C,将其展平为一维向量,然后通过一个线性层(全连接层)投影到维度为 DDD 的向量空间。这个 DDD 就是 Transformer 的隐藏层维度(Hidden Size)。
在代码实现中,Patch Embedding 通常用一个卷积核大小为 PPP、步长为 PPP 的卷积层来实现,这样既高效又简洁。
3.2 位置编码
Transformer 本身是置换不变的(Permutation Invariant),即打乱输入 Token 的顺序,输出结果不变。但图像 Patch 的顺序是有意义的,因此需要为每个 Patch 添加位置编码。
ViT 使用可学习的位置编码(Learnable Position Embedding),即初始化一组与 Patch 数量相同的可学习向量,在训练过程中自动学习每个位置合适的编码。位置编码与 Patch Embedding 直接相加,得到最终的输入向量。
3.3 [CLS] Token
在 BERT 中,[CLS] Token 被添加到序列开头,用于聚合整个序列的信息。ViT 沿用了这一设计:在 Patch 序列的最前面添加一个可学习的 [CLS] Token,它与所有 Patch 一起参与自注意力计算。
经过 Transformer 编码器后,[CLS] Token 对应的输出向量就包含了整个图像的全局信息,可以接一个分类头用于图像分类。在 VLA 模型中,这个 [CLS] Token 的输出(或所有 Patch 的输出)会被送入语言模型作为视觉特征。
3.4 Transformer Encoder
Transformer Encoder 由多个相同的 Block 堆叠而成,每个 Block 包含两个核心子层:
- 多头自注意力(Multi-Head Self-Attention, MSA):让每个 Token 与其他所有 Token 交互,建模全局依赖关系。
- 多层感知机(MLP):逐位置的前馈网络,通常包含两个全连接层和 GELU 激活函数。
每个子层都使用了残差连接(Residual Connection)和层归一化(Layer Normalization, LN)。ViT 采用的是 Pre-LN 结构,即先做 LayerNorm 再做注意力/MLP。
单个 Transformer Block 的内部计算流程如下:
3.5 分类头
对于图像分类任务,ViT 将 [CLS] Token 的输出经过 LayerNorm 后,送入一个线性分类头,输出各类别的 logits。在 VLA 模型中,分类头会被替换为与语言模型对齐的投影层。
下图展示了 ViT 的完整模型结构:
4. ViT 的数学原理
4.1 自注意力机制
自注意力是 Transformer 的核心。对于输入序列 X∈RN×DX \in \mathbb{R}^{N \times D}X∈RN×D,通过三个可学习的权重矩阵 WQ,WK,WVW_Q, W_K, W_VWQ,WK,WV 计算 Query、Key、Value:
Q=XWQ,K=XWK,V=XWVQ = XW_Q, \quad K = XW_K, \quad V = XW_VQ=XWQ,K=XWK,V=XWV
然后计算注意力分数:
Attention(Q,K,V)=softmax(QKTdk)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)VAttention(Q,K,V)=softmax(dkQKT)V
其中 dkd_kdk 是 Query/Key 的维度,除以 dk\sqrt{d_k}dk 是为了防止点积结果过大导致梯度消失。
4.2 多头注意力
多头注意力将 DDD 维的 Query、Key、Value 拆分为 hhh 个头,每个头的维度为 dk=D/hd_k = D / hdk=D/h。每个头独立计算注意力,然后将结果拼接并线性投影:
MSA(X)=Concat(head1,…,headh)WO\text{MSA}(X) = \text{Concat}(\text{head}_1, \ldots, \text{head}_h)W_OMSA(X)=Concat(head1,…,headh)WO
多头机制允许模型在不同的表示子空间中关注不同的信息,增强了模型的表达能力。
4.3 MLP 与残差连接
每个 Transformer Block 的计算过程可以表示为:
X′=X+MSA(LN(X))X' = X + \text{MSA}(\text{LN}(X))X′=X+MSA(LN(X))
X′′=X′+MLP(LN(X′))X'' = X' + \text{MLP}(\text{LN}(X'))X′′=X′+MLP(LN(X′))
其中 MLP 通常为:
MLP(Z)=GELU(ZW1+b1)W2+b2\text{MLP}(Z) = \text{GELU}(ZW_1 + b_1)W_2 + b_2MLP(Z)=GELU(ZW1+b1)W2+b2
GELU(Gaussian Error Linear Unit)是 ViT 中使用的激活函数,可以看作 ReLU 的平滑版本。
自注意力的计算过程可以用下图表示:
5. ViT 的 PyTorch 复现
下面使用 PyTorch 从零开始实现一个简化版的 ViT。为了便于理解,采用模块化的方式逐步实现。
5.1 导入依赖库
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
from typing import Optional, Tuple
5.2 Patch Embedding 实现
class PatchEmbedding(nn.Module):
"""将图像切分为 Patch 并投影为向量。
使用卷积核大小 = 步长 = patch_size 的卷积实现,
等价于先切 Patch 再展平后做线性投影。
"""
def __init__(self, in_channels: int = 3, patch_size: int = 16,
embed_dim: int = 768):
super().__init__()
self.patch_size = patch_size
self.proj = nn.Conv2d(
in_channels,
embed_dim,
kernel_size=patch_size,
stride=patch_size,
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: [B, C, H, W]
B, C, H, W = x.shape
assert H % self.patch_size == 0 and W % self.patch_size == 0, \
f"图像尺寸 {H}x{W} 必须能被 patch_size={self.patch_size} 整除"
x = self.proj(x) # [B, embed_dim, H/P, W/P]
x = x.flatten(2) # [B, embed_dim, N]
x = x.transpose(1, 2) # [B, N, embed_dim]
return x
5.3 多头自注意力实现
class MultiHeadSelfAttention(nn.Module):
"""多头自注意力机制。"""
def __init__(self, embed_dim: int, num_heads: int,
dropout: float = 0.0):
super().__init__()
assert embed_dim % num_heads == 0, \
f"embed_dim={embed_dim} 必须能被 num_heads={num_heads} 整除"
self.embed_dim = embed_dim
self.num_heads = num_heads
self.head_dim = embed_dim // num_heads
self.scale = self.head_dim ** -0.5
self.qkv = nn.Linear(embed_dim, embed_dim * 3)
self.proj = nn.Linear(embed_dim, embed_dim)
self.attn_drop = nn.Dropout(dropout)
self.proj_drop = nn.Dropout(dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
B, N, D = x.shape
# 计算 Q、K、V
qkv = self.qkv(x) # [B, N, 3*D]
qkv = qkv.reshape(B, N, 3, self.num_heads, self.head_dim)
qkv = qkv.permute(2, 0, 3, 1, 4) # [3, B, num_heads, N, head_dim]
q, k, v = qkv.unbind(0)
# 计算注意力分数
attn = (q @ k.transpose(-2, -1)) * self.scale # [B, num_heads, N, N]
attn = attn.softmax(dim=-1)
attn = self.attn_drop(attn)
# 加权求和
x = (attn @ v) # [B, num_heads, N, head_dim]
x = x.transpose(1, 2).reshape(B, N, D) # [B, N, D]
x = self.proj(x)
x = self.proj_drop(x)
return x
5.4 MLP 实现
class MLP(nn.Module):
"""前馈网络:两个全连接层 + GELU 激活。"""
def __init__(self, embed_dim: int, mlp_ratio: float = 4.0,
dropout: float = 0.0):
super().__init__()
hidden_dim = int(embed_dim * mlp_ratio)
self.fc1 = nn.Linear(embed_dim, hidden_dim)
self.fc2 = nn.Linear(hidden_dim, embed_dim)
self.act = nn.GELU()
self.drop = nn.Dropout(dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.fc1(x)
x = self.act(x)
x = self.drop(x)
x = self.fc2(x)
x = self.drop(x)
return x
5.5 Transformer Block 实现
class TransformerBlock(nn.Module):
"""单个 Transformer Block:Pre-LN 结构。"""
def __init__(self, embed_dim: int, num_heads: int,
mlp_ratio: float = 4.0, dropout: float = 0.0):
super().__init__()
self.norm1 = nn.LayerNorm(embed_dim)
self.attn = MultiHeadSelfAttention(embed_dim, num_heads, dropout)
self.norm2 = nn.LayerNorm(embed_dim)
self.mlp = MLP(embed_dim, mlp_ratio, dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = x + self.attn(self.norm1(x))
x = x + self.mlp(self.norm2(x))
return x
5.6 完整 ViT 模型实现
class ViT(nn.Module):
"""Vision Transformer 完整实现。
参数:
img_size: 输入图像尺寸(假设为正方形)
patch_size: Patch 尺寸
in_channels: 输入图像通道数
num_classes: 分类类别数
embed_dim: 隐藏层维度
depth: Transformer Block 数量
num_heads: 注意力头数
mlp_ratio: MLP 隐藏层维度倍数
dropout: Dropout 概率
"""
def __init__(self, img_size: int = 224, patch_size: int = 16,
in_channels: int = 3, num_classes: int = 1000,
embed_dim: int = 768, depth: int = 12,
num_heads: int = 12, mlp_ratio: float = 4.0,
dropout: float = 0.0):
super().__init__()
self.patch_embed = PatchEmbedding(in_channels, patch_size, embed_dim)
num_patches = (img_size // patch_size) ** 2
self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
self.pos_embed = nn.Parameter(
torch.zeros(1, num_patches + 1, embed_dim)
)
self.pos_drop = nn.Dropout(dropout)
self.blocks = nn.Sequential(*[
TransformerBlock(embed_dim, num_heads, mlp_ratio, dropout)
for _ in range(depth)
])
self.norm = nn.LayerNorm(embed_dim)
self.head = nn.Linear(embed_dim, num_classes)
# 初始化权重
nn.init.trunc_normal_(self.pos_embed, std=0.02)
nn.init.trunc_normal_(self.cls_token, std=0.02)
self.apply(self._init_weights)
def _init_weights(self, module):
if isinstance(module, nn.Linear):
nn.init.trunc_normal_(module.weight, std=0.02)
if module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.LayerNorm):
nn.init.ones_(module.weight)
nn.init.zeros_(module.bias)
def forward(self, x: torch.Tensor) -> torch.Tensor:
B = x.shape[0]
# Patch Embedding
x = self.patch_embed(x) # [B, N, D]
# 拼接 [CLS] Token
cls_token = self.cls_token.expand(B, -1, -1)
x = torch.cat([cls_token, x], dim=1) # [B, N+1, D]
# 添加位置编码
x = x + self.pos_embed
x = self.pos_drop(x)
# Transformer Encoder
x = self.blocks(x)
# 取 [CLS] Token 输出
x = self.norm(x[:, 0])
# 分类头
x = self.head(x)
return x
5.7 测试模型
def test_vit():
"""测试 ViT 模型的前向传播。"""
# 使用小尺寸参数便于快速验证
model = ViT(
img_size=32,
patch_size=8,
in_channels=3,
num_classes=10,
embed_dim=128,
depth=4,
num_heads=4,
mlp_ratio=2.0,
dropout=0.1,
)
# 随机生成一个 batch 的输入
x = torch.randn(2, 3, 32, 32)
out = model(x)
print(f"输入形状: {x.shape}")
print(f"输出形状: {out.shape}")
assert out.shape == (2, 10), f"输出形状错误: {out.shape}"
print("✅ 前向传播测试通过!")
if __name__ == "__main__":
test_vit()
运行上述代码,输出结果如下:
输入形状: torch.Size([2, 3, 32, 32])
输出形状: torch.Size([2, 10])
✅ 前向传播测试通过!
在动手实现之前,先用一张图梳理各模块之间的依赖关系与数据流向:
6. 在 VLA 中使用 ViT
在 VLA 模型中,ViT 通常作为视觉编码器,将机器人摄像头采集的图像转换为视觉特征序列。与图像分类任务不同,VLA 中的 ViT 有以下几点差异:
- 不使用分类头:VLA 不需要输出类别概率,而是将 [CLS] Token 或所有 Patch Token 的输出作为视觉特征,送入后续的语言模型。
- 分辨率更高:机器人观测的图像通常分辨率较高(如 224×224 或更高),以保留足够的空间细节。
- 与语言模型对齐:ViT 输出的特征维度需要与语言模型的输入维度对齐,通常通过一个投影层实现。
- 可能使用冻结或微调策略:在 VLA 训练中,ViT 可以冻结(不更新参数)或进行微调,取决于数据量和训练策略。
下面是一个在 VLA 中使用 ViT 的简化示例:
class VLAVisionEncoder(nn.Module):
"""VLA 中的视觉编码器:ViT + 投影层。"""
def __init__(self, img_size: int = 224, patch_size: int = 16,
embed_dim: int = 768, depth: int = 12,
num_heads: int = 12, llm_dim: int = 4096):
super().__init__()
# 加载预训练 ViT(这里用我们上面实现的 ViT)
self.vit = ViT(
img_size=img_size,
patch_size=patch_size,
num_classes=0, # 不需要分类头
embed_dim=embed_dim,
depth=depth,
num_heads=num_heads,
)
# 移除分类头
self.vit.head = nn.Identity()
# 投影层:将 ViT 输出维度映射到 LLM 输入维度
self.proj = nn.Linear(embed_dim, llm_dim)
def forward(self, images: torch.Tensor) -> torch.Tensor:
# images: [B, C, H, W]
B = images.shape[0]
# 提取视觉特征
x = self.vit.patch_embed(images) # [B, N, D]
# 拼接 [CLS] Token 并添加位置编码
cls_token = self.vit.cls_token.expand(B, -1, -1)
x = torch.cat([cls_token, x], dim=1)
x = x + self.vit.pos_embed
# Transformer Encoder
x = self.vit.blocks(x)
x = self.vit.norm(x)
# 投影到 LLM 维度
x = self.proj(x) # [B, N+1, llm_dim]
return x
最后,用一张图总结 ViT 在 VLA 模型中的位置与作用:
本文深入学习了 ViT 的算法原理,并动手用 PyTorch 从零实现了一个完整的 ViT 模型。回顾一下本文的关键知识点:
今天我们深入学习了 ViT 的算法原理,并动手用 PyTorch 从零复现了一个完整的 ViT 模型。让我们回顾一下今天的关键知识点:
- ViT 的核心思想:将图像切分为 Patch,线性投影为向量,输入标准 Transformer 编码器。
- 模型结构:Patch Embedding + 位置编码 + [CLS] Token + Transformer Encoder + 分类头。
- 数学原理:自注意力机制、多头注意力、残差连接与 LayerNorm。
- 代码实现:从 Patch Embedding 到完整 ViT 的模块化 PyTorch 实现。
- VLA 中的应用:ViT 作为视觉编码器,输出视觉特征供语言模型使用。
明天我们将继续学习 VLA 中的视觉-语言对齐机制,了解如何将 ViT 提取的视觉特征与文本特征进行对齐,敬请期待!
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐


所有评论(0)