公司动态
从注意力公式到ViT:Transformer核心机制与视觉应用全解析
1. 项目概述从“注意力”到“视觉革命”的认知地图如果你最近在接触深度学习尤其是自然语言处理NLP或计算机视觉CV那么“Transformer”、“Self-Attention”、“ViT”这几个词一定像潮水一样涌现在你眼前。它们不再是某个遥远论文里的晦涩概念而是构建当今绝大多数AI大模型的基石。我第一次系统学习Transformer时面对那一堆公式和矩阵运算感觉就像在看天书尤其是那个核心的“注意力公式”每一步都知其然不知其所以然。后来在复现Vision TransformerViT时更是被其中模块的拼接逻辑搞得晕头转向。这促使我花了大量时间从一个实践者的角度去拆解、消化和重建对这些概念的理解。今天我想和你分享的不是一篇照本宣科的论文翻译而是一份结合了代码、图示和大量“踩坑”经验的认知地图。我们将彻底搞懂那个看似复杂的注意力公式每一步到底在计算什么、为什么要这么计算然后深入Self-Attention自注意力机制看它如何让模型学会“抓重点”最后我们将把所有这些知识应用到Vision TransformerViT上完整理解一个图像是如何被Transformer“看懂”的。无论你是刚入门的新手还是想巩固理解的从业者这篇文章都将提供一条清晰、可操作的路径让你不仅能理解原理更能想象出代码背后的逻辑甚至能自己动手复现。让我们暂时忘掉那些令人望而生畏的数学符号从最直观的“为什么”开始。2. 注意力公式的逐层拆解不只是矩阵乘法很多人一看到注意力公式就头疼觉得是复杂的数学游戏。但如果我们把它还原成一个“信息检索”的过程一切就清晰了。想象一下你正在阅读一篇长文当看到“苹果”这个词时你的大脑会瞬间关联到上下文是“水果”还是“公司”。注意力机制干的就是这个事对于序列中的每一个元素比如一个词它要决定应该“注意”序列中其他哪些元素的信息。公式是这个过程的高度数学抽象。2.1 核心公式再现与全局视角我们首先回顾一下最经典的缩放点积注意力Scaled Dot-Product Attention公式Attention(Q, K, V) softmax( (QK^T) / √d_k ) V这个公式里有三个核心输入查询Query、键Key和值Value通常由输入序列通过不同的线性变换得到。输出是一个加权求和后的新表示。下面我们一步步拆解。2.2 第一步Q与K的点积 (QK^T) —— 计算相关性分数操作将查询矩阵Q和键矩阵K的转置进行矩阵乘法。含义这一步的目的是计算输入序列中每一个元素作为Query与序列中所有元素包括自己作为Key之间的相关性或相似度。为什么是点积在向量空间中两个向量的点积可以衡量它们的相似程度。点积值越大通常意味着两个向量方向越接近相关性越高。这里每个Query向量会与所有Key向量做点积得到一个分数矩阵其中第i行第j列的元素就表示第i个Query与第j个Key的关联强度。实操细节假设输入序列有n个元素每个元素的特征维度是d_model。经过线性投影后Q、K的维度通常是[n, d_k]这里d_k是投影后的维度。那么QK^T的维度就是[n, n]。这个n x n的矩阵就是原始的注意力分数矩阵它尚未经过归一化和缩放。注意这里有一个非常关键的“为什么”。为什么用点积而不是余弦相似度虽然余弦相似度也衡量方向但点积同时考虑了向量的长度模长。在训练初期模型参数是随机初始化的点积的绝对值可能非常大或非常小这会导致后续softmax函数的梯度问题进入饱和区。这就是下一步缩放存在的根本原因。2.3 第二步缩放 (除以√d_k) —— 稳定梯度流动操作将上一步得到的分数矩阵中的每一个元素都除以一个缩放因子 √d_kd_k是Key向量的维度。含义对注意力分数进行缩放使其数值分布更加稳定。为什么需要缩放这是整个公式中最精妙的设计之一直接关系到训练的稳定性。点积的结果的方差会随着向量维度d_k的增大而增大。假设Q和K中的元素是均值为0、方差为1的独立随机变量那么点积q·k的均值是0方差就是d_k。当d_k很大时例如512或1024点积值的绝对值可能会变得非常大。而后续的softmax函数对输入绝对值非常敏感过大的输入会使softmax的输出无限接近一个one-hot向量即概率几乎全集中在一个位置上这会导致梯度变得非常小梯度消失使得模型难以更新学习。除以√d_k可以将点积的方差重新缩放回1左右从而使得softmax函数的输入保持在梯度敏感的区域保障了训练的有效性。实操心得在你自己实现注意力层时这个缩放步骤千万不能省略。我曾在早期实验中忘记除以√d_k结果模型损失完全不下降注意力权重几乎变成了随机噪声排查了很久才找到这个原因。这是一个经典的“坑”。2.4 第三步Softmax归一化 —— 产生注意力权重操作对缩放后的分数矩阵的每一行应用softmax函数。含义将每一行的分数即一个Query对所有Key的相似度转换为一个概率分布所有值在0到1之间且和为1。为什么是SoftmaxSoftmax函数具有“放大优胜者”的特性。它会让最高的分数对应的概率更高而抑制较低的分数。这模拟了人类注意力“聚焦”的特点——我们更关注最相关的信息而不是平均分配注意力。经过softmax后得到的矩阵我们称之为“注意力权重矩阵”Attention Weight Matrix它的形状同样是[n, n]。第i行就代表了第i个位置对于序列所有位置的关注程度权重。细节与技巧在代码实现中特别是使用PyTorch或TensorFlow时直接调用F.softmax()或tf.nn.softmax时通常指定dim-1即对最后一个维度每一行进行归一化。此外为了数值稳定性在实际计算中缩放和softmax通常会合并处理避免中间值过大。2.5 第四步与V加权求和 (乘以V) —— 聚合信息操作将上一步得到的注意力权重矩阵与值矩阵V进行矩阵乘法。含义根据计算出的注意力权重对值Value信息进行加权求和得到每个位置的最终输出表示。为什么是VQuery和Key的作用是计算“谁更重要”而Value才是真正承载了要被提取和聚合的信息本体。你可以把Key看作信息的“索引”或“标签”而Value是“内容”。通过注意力权重模型决定从每个位置的“内容”Value中抽取多少混合起来形成当前Query位置的新表示。这个新表示既包含了自身的信息也以加权的方式融入了全局上下文中最相关的信息。输出理解假设V的维度是[n, d_v]那么输出矩阵的维度就是[n, d_v]。这意味着序列中每个位置都得到了一个新的、融合了全局上下文信息的d_v维向量表示。至此我们完成了对单头注意力公式的微观拆解。每一个步骤都有其明确的物理意义和数学必要性绝非随意设计。理解了这个就握住了打开Transformer世界大门的钥匙。3. Self-Attention自注意力机制让模型学会“抓重点”理解了基础的注意力计算Self-Attention就水到渠成了。它的核心思想简单而强大让序列中的每个元素通过注意力机制与序列中的所有元素包括自己进行交互从而动态地计算出一个新的、富含上下文信息的表示。3.1 Self-Attention的核心思想与计算图定义在Self-Attention中Query、Key、Value三者都来源于同一个输入序列。也就是说模型自己对自己做注意力。输入一个序列X形状[n, d_model]我们通过三个不同的可学习权重矩阵W_Q,W_K,W_V将其分别投影到Q, K, V空间Q X * W_QK X * W_KV X * W_V然后将得到的Q, K, V代入上一章讲解的注意力公式进行计算。这个过程允许序列中的任意两个位置直接建立联系无论它们相距多远。这与RNN/LSTM需要逐步传递信息形成鲜明对比后者在长距离依赖上存在梯度消失或爆炸的问题。计算流程可视化输入X [n, d_model] ↓ (线性变换 W_Q) 查询Q [n, d_k] ↓ |---- (Q * K^T) / √d_k - Softmax - 注意力权重 [n, n] | 键K [n, d_k] ---↑ | 值V [n, d_v] ------------------- 加权求和 ↓ 输出O [n, d_v]3.2 为什么Self-Attention如此强大并行计算能力注意力权重矩阵的计算QK^T是纯粹的矩阵乘法可以高度并行化充分利用GPU等硬件加速。这与RNN的序列依赖计算有本质区别也是Transformer训练速度快的根本原因。全局感知野在计算输出时每个位置都能直接“看到”序列中的所有其他位置并从中提取信息。这使得模型能够轻松捕获长距离的依赖关系。动态权重注意力权重不是静态的而是根据具体的输入内容动态计算的。对于不同的输入句子“苹果”与“吃”和“股价”的关联强度是不同的。这种动态性赋予了模型强大的内容理解能力。可解释性通过可视化注意力权重矩阵我们能够直观地看到模型在做出决策时“关注”了输入序列的哪些部分。这为理解模型行为提供了一扇窗口。3.3 从单头到多头注意力 (Multi-Head Attention)单头注意力就像只用一种“视角”或“理解方式”去分析句子。而多头注意力则引入了多个这样的视角让模型能够同时关注来自不同表示子空间的信息。操作步骤线性投影到多子空间对输入X我们使用h组h个头不同的(W_Q, W_K, W_V)投影矩阵分别得到h组(Q_i, K_i, V_i)。并行计算多头注意力对每一组(Q_i, K_i, V_i)独立进行上一章所述的缩放点积注意力计算得到h个输出头head_i形状为[n, d_v]通常令d_v d_model / h以保持总参数量不变。拼接与最终投影将h个head_i在特征维度上拼接起来得到一个形状为[n, h * d_v]即[n, d_model]的大矩阵。最后通过一个可学习的输出投影矩阵W_O形状[d_model, d_model]进行线性变换得到最终的Multi-Head Attention输出。公式表示MultiHead(Q, K, V) Concat(head_1, ..., head_h) * W_Owhere head_i Attention(Q * W_Q_i, K * W_K_i, V * W_V_i)为什么需要多头这类似于卷积神经网络中使用多个滤波器。不同的头可以学习到不同的注意力模式。例如在处理一个句子时一个头可能主要关注语法结构如主谓一致另一个头可能关注指代关系如代词指向哪个名词再一个头可能关注语义关联。通过多头的并行计算与融合模型的理解能力变得更加丰富和鲁棒。实操心得在实现多头注意力时一个高效的技巧是使用矩阵运算的“批处理”思想。我们并不需要真的用for循环去计算h次。而是将W_Q, W_K, W_V的维度设计为[d_model, h, d_k]然后通过einsum或reshape操作一次性地计算出所有头的Q、K、V再通过调整维度利用矩阵乘法一次性完成所有头的注意力计算。这能极大提升计算效率。PyTorch中的nn.MultiheadAttention模块就采用了这种实现。4. Vision Transformer (ViT) 模块全解当Transformer遇见图像理解了Self-Attention和多头注意力我们就可以进军计算机视觉领域看看Transformer是如何处理图像的。Vision Transformer (ViT) 是这一领域的开创性工作它的核心思想非常直接将图像视为一个序列的图块Patch然后直接用标准的Transformer Encoder来处理这个序列。4.1 ViT的整体架构与数据流ViT抛弃了CNN中固有的归纳偏置如局部性、平移不变性完全依赖注意力机制来学习图像特征。其架构可以清晰地分为以下几个步骤图块嵌入 (Patch Embedding)输入一张图像X形状为[H, W, C]高、宽、通道数如224x224x3。操作将图像分割成固定大小的非重叠图块。例如使用16x16的图块大小那么一张224x224的图像会被分成 (224/16) * (224/16) 14 * 14 196个图块。实现每个图块形状[16, 16, 3] 768维被展平为一个向量然后通过一个可训练的线性投影层全连接层映射到Transformer的隐藏维度D例如768。这个线性投影层就起到了“嵌入”的作用将原始的像素空间映射到模型语义空间。输出是一个形状为[196, 768]的序列我们可以将其视为196个“视觉词”的序列。位置编码 (Position Embedding)问题Self-Attention本身是置换不变的打乱输入序列顺序输出不变。但图像中图块的空间位置信息至关重要。解决方案为序列中的每一个位置即每一个图块添加一个可学习的位置编码向量。这个向量的维度与图块嵌入的维度相同768。操作生成一个形状为[197, 768]的可学习参数矩阵为什么是197见下一步将其加到图块嵌入序列上。这样模型在计算注意力时就能感知到每个图块在原始图像中的绝对或相对位置信息。分类令牌 (Class Token)借鉴自BERT在序列的最前面额外添加一个可学习的嵌入向量称为[class] token。这个令牌本身不对应任何图像图块。作用在通过Transformer Encoder进行信息交互后这个[class] token的最终状态汇聚了整个图像的信息被用来作为整个图像的表示送入最后的分类头MLP进行类别预测。因此最终的输入序列长度变为196图块 1[class] token 197。Transformer Encoder 堆叠结构输入序列197个长度为768的向量经过L个相同的Transformer Encoder层。单层Encoder构成这是ViT的核心模块每一层都包含两个主要子层 a.多头自注意力层 (Multi-Head Self-Attention, MSA)就是我们之前详解的部分。序列内部所有元素包括[class] token和图块进行全局信息交互。 b.前馈网络层 (Feed-Forward Network, FFN)通常是一个两层MLP对每个位置的向量进行独立的、非线性的变换。常见配置是第一层将维度从D扩大到4D如768-3072使用GELU激活函数第二层再投影回D维度3072-768。残差连接与层归一化每个子层MSA和FFN都应用了残差连接和层归一化LayerNorm。具体顺序是LayerNorm - 子层计算 - Add残差。这种“Pre-Norm”结构有助于训练的稳定性。公式化表示对于第l层输入为z_{l-1}z_l MSA(LayerNorm(z_{l-1})) z_{l-1}z_l FFN(LayerNorm(z_l)) z_lMLP 分类头输入取最后一个Transformer Encoder层输出的[class] token对应的向量形状[1, D]。结构通常是一个简单的多层感知机可能包含一个隐藏层和Dropout最后输出对应分类类别的logits。4.2 ViT中的自注意力图像上下文的理解在ViT中自注意力机制是如何工作的呢对于那197个令牌1个[class] 196个图块组成的序列多头自注意力层会让它们两两之间计算注意力。[class] token的视角这个特殊的令牌会关注所有图像图块。通过训练它学会从全局“收集”与分类任务最相关的信息。最终它的表示就承载了整张图像的语义摘要。图块之间的视角一个描绘“狗耳朵”的图块可能会高度关注“狗鼻子”、“狗眼睛”的图块从而在特征表示上强化“狗头”这一局部概念。同时它也可能微弱地关注到远处的“狗尾巴”图块建立起整体的物体结构信息。这种全局的、动态的关联能力是传统CNN通过堆叠卷积层逐步扩大感受野所难以直接实现的。4.3 与CNN的对比及ViT的优劣优势全局建模能力从第一层开始就具备全局感受野能直接建模图像中任意两个区域的长距离依赖。可扩展性模型性能随着数据量、模型规模深度、宽度、参数量的增加而显著提升展现出强大的缩放定律Scaling Law。架构统一与NLP领域的Transformer架构统一便于进行多模态研究和模型迁移。挑战与劣势数据饥渴ViT缺乏CNN固有的图像归纳偏置因此在小规模数据集如ImageNet-1K上从头训练时效果通常不如经过高度优化的CNN如ResNet。它需要在大规模数据集如JFT-300M上预训练才能发挥全部潜力。计算复杂度自注意力机制的计算复杂度与序列长度的平方成正比O(n²)。对于高分辨率图像图块数量n会很大导致计算和内存开销急剧上升。后续的Swin Transformer等工作通过引入局部窗口注意力、移位窗口等机制有效地缓解了这个问题。细节信息将图像分割成较大图块如16x16可能会损失一些细粒度的局部信息。5. 从理解到实践核心代码实现与调试要点理论理解了不落实到代码总是感觉不踏实。这里我用PyTorch风格伪代码勾勒出最核心模块的实现并附上关键调试经验。5.1 缩放点积注意力实现import torch import torch.nn.functional as F def scaled_dot_product_attention(q, k, v, maskNone): Args: q: Query tensor, shape [batch_size, seq_len_q, d_k] k: Key tensor, shape [batch_size, seq_len_k, d_k] v: Value tensor, shape [batch_size, seq_len_k, d_v] mask: Optional mask tensor, shape [batch_size, seq_len_q, seq_len_k] Returns: output: weighted value tensor, shape [batch_size, seq_len_q, d_v] attention_weights: softmax后的权重 shape [batch_size, seq_len_q, seq_len_k] # 1. 计算点积分数 matmul_qk torch.matmul(q, k.transpose(-2, -1)) # [..., seq_len_q, seq_len_k] # 2. 缩放 d_k q.size(-1) scaled_attention_logits matmul_qk / (d_k ** 0.5) # 3. 可选应用掩码如解码器的因果掩码 if mask is not None: # 将mask中为1的位置需要被掩盖替换为一个非常大的负数使得softmax后概率为0 scaled_attention_logits scaled_attention_logits.masked_fill(mask 0, -1e9) # 4. Softmax归一化得到权重 attention_weights F.softmax(scaled_attention_logits, dim-1) # 5. 加权求和 output torch.matmul(attention_weights, v) # [..., seq_len_q, d_v] return output, attention_weights5.2 多头注意力层实现class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads): super().__init__() assert d_model % num_heads 0, d_model must be divisible by num_heads self.d_model d_model self.num_heads num_heads self.d_k d_model // num_heads # 定义投影矩阵注意这里将多个头的参数合并到了一个大的线性层中 self.w_q nn.Linear(d_model, d_model) # 输出维度是 d_model num_heads * d_k self.w_k nn.Linear(d_model, d_model) self.w_v nn.Linear(d_model, d_model) self.w_o nn.Linear(d_model, d_model) # 输出投影层 def split_heads(self, x, batch_size): 将最后的d_model维度分割成 (num_heads, d_k) x x.view(batch_size, -1, self.num_heads, self.d_k) # 调整维度为 [batch_size, num_heads, seq_len, d_k] 以便并行计算 return x.permute(0, 2, 1, 3) def forward(self, q, k, v, maskNone): batch_size q.size(0) # 1. 线性投影并分割头 q self.split_heads(self.w_q(q), batch_size) # [B, h, seq_len_q, d_k] k self.split_heads(self.w_k(k), batch_size) # [B, h, seq_len_k, d_k] v self.split_heads(self.w_v(v), batch_size) # [B, h, seq_len_k, d_v] (d_v d_k) # 2. 使用缩放点积注意力计算每个头 # 注意这里attention函数的输入维度是 [batch_size*num_heads, seq_len, d_k] # 我们需要将头维度与批次维度合并以实现并行计算 q q.reshape(batch_size * self.num_heads, -1, self.d_k) k k.reshape(batch_size * self.num_heads, -1, self.d_k) v v.reshape(batch_size * self.num_heads, -1, self.d_k) if mask is not None: # mask需要广播到所有头 mask mask.unsqueeze(1) # [B, 1, seq_len_q, seq_len_k] mask mask.repeat(1, self.num_heads, 1, 1) # [B, h, seq_len_q, seq_len_k] mask mask.reshape(batch_size * self.num_heads, -1, mask.size(-1)) scaled_attention, attention_weights scaled_dot_product_attention(q, k, v, mask) # 3. 合并头 scaled_attention scaled_attention.reshape(batch_size, self.num_heads, -1, self.d_k) scaled_attention scaled_attention.permute(0, 2, 1, 3) # [B, seq_len_q, h, d_k] concat_attention scaled_attention.reshape(batch_size, -1, self.d_model) # [B, seq_len_q, d_model] # 4. 最终线性投影 output self.w_o(concat_attention) return output, attention_weights5.3 ViT图块嵌入与位置编码实现class PatchEmbedding(nn.Module): 将图像分割为图块并嵌入到向量空间 def __init__(self, img_size224, patch_size16, in_channels3, embed_dim768): super().__init__() self.img_size (img_size, img_size) self.patch_size (patch_size, patch_size) self.num_patches (img_size // patch_size) ** 2 # 使用一个卷积层来实现图块展平和投影效率更高 # 卷积核大小和步长等于图块大小输出通道数等于嵌入维度 self.projection nn.Conv2d(in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): # x: [B, C, H, W] B, C, H, W x.shape assert H self.img_size[0] and W self.img_size[1], \ fInput image size ({H}*{W}) doesnt match model ({self.img_size[0]}*{self.img_size[1]}). x self.projection(x) # [B, embed_dim, num_patches_h, num_patches_w] x x.flatten(2) # 将高和宽维度展平 - [B, embed_dim, num_patches] x x.transpose(1, 2) # [B, num_patches, embed_dim] return x class VisionTransformer(nn.Module): def __init__(self, img_size224, patch_size16, in_channels3, num_classes1000, embed_dim768, depth12, num_heads12, mlp_ratio4.0): super().__init__() self.patch_embed PatchEmbedding(img_size, patch_size, in_channels, embed_dim) num_patches self.patch_embed.num_patches # [class] token 和位置编码 self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed nn.Parameter(torch.zeros(1, num_patches 1, embed_dim)) # 1 for cls_token self.pos_drop nn.Dropout(p0.1) # Transformer Encoder 层堆叠 self.blocks nn.ModuleList([ TransformerBlock(embed_dim, num_heads, mlp_ratio) for _ in range(depth) ]) self.norm nn.LayerNorm(embed_dim) # 分类头 self.head nn.Linear(embed_dim, num_classes) # 初始化权重 self._init_weights() def _init_weights(self): # 简单起见使用截断正态分布初始化 nn.init.trunc_normal_(self.pos_embed, std0.02) nn.init.trunc_normal_(self.cls_token, std0.02) # 其他线性层和LayerNorm有默认初始化通常足够 def forward(self, x): B x.shape[0] # 1. 图块嵌入 x self.patch_embed(x) # [B, num_patches, embed_dim] # 2. 添加 [class] token cls_tokens self.cls_token.expand(B, -1, -1) # 从 [1,1,D] 扩展到 [B,1,D] x torch.cat((cls_tokens, x), dim1) # [B, num_patches1, embed_dim] # 3. 添加位置编码 x x self.pos_embed x self.pos_drop(x) # 4. 通过Transformer Encoder for blk in self.blocks: x blk(x) # 5. 取 [class] token 的状态用于分类 x self.norm(x) cls_output x[:, 0] # 取第一个位置即[class] token的输出 # 6. 分类头 logits self.head(cls_output) return logits5.4 关键调试经验与“避坑”指南维度对齐是生命线在实现注意力或多头注意力时最常出现的bug就是维度不匹配。务必清楚每一步操作后张量的形状。善用print(x.shape)或调试器来跟踪维度变化。特别注意transpose,permute,view,reshape这些操作的区别和联系。注意力掩码的应用在训练语言模型或图像生成任务时经常需要因果掩码防止看到未来信息或填充掩码忽略padding部分。确保掩码在softmax之前正确应用并且形状能广播到注意力分数矩阵。梯度消失/爆炸与初始化Transformer模型较深初始化很重要。对于线性层常用的初始化有nn.init.xavier_uniform_或nn.init.kaiming_uniform_。对于位置编码ViT使用截断正态分布。LayerNorm是稳定训练的关键组件。学习率与优化器AdamW优化器配合热身Warmup学习率调度是训练Transformer的黄金标准。Warmup阶段例如前5000步让学习率从0线性增加到预设值有助于模型在训练初期稳定。可视化注意力权重这是理解模型在“看”哪里的强大工具。你可以提取中间某层的注意力权重矩阵attention_weights将其对应回原始图像图块用热力图的形式叠加在图像上。这不仅能帮你调试还能增加对模型行为的信任。ViT训练的数据增强由于ViT数据饥渴在中小数据集上训练时强力的数据增强如RandAugment, MixUp, CutMix至关重要。这相当于人为增加了数据的多样性弥补了模型本身归纳偏置的不足。混合精度训练使用torch.cuda.amp进行自动混合精度训练可以显著减少GPU内存占用并加快训练速度尤其对于ViT这类大模型。但要注意在计算softmax或损失函数时可能需要保持fp32精度以避免数值下溢。理解这些模块的代码实现再结合理论你就能真正掌握从公式到可运行模型的全过程。下次当你看到或使用Transformer相关的代码时就不会再觉得它是一个黑盒而能清晰地看到数据流动的每一个环节。