ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

ViT论文精读:图像切成16x16个词,Transformer如何颠覆CNN

ViT论文精读:图像切成16x16个词,Transformer如何颠覆CNN ViT 这篇论文圈内习惯叫它《An Image is Worth 16x16 Words》自从 2020 年底放出来之后直接把计算机视觉的研究方向带偏了一大截——不对是带进了一个新阶段。以前大家一说视觉模型就是 CNN 的天下ResNet、EfficientNet 各种卷积结构卷得不亦乐乎ViT 出来之后大家才意识到原来把图片切成一块一块的 patch然后当成词向量丢给 Transformer也能在图像分类上干翻 CNN而且预训练数据量越大优势越明显。这篇论文我前前后后精读了很多遍也做过完整的代码复现和翻译注释今天把翻译和拆解一起整理出来希望能帮你省掉啃原论文的力气。这篇内容适合谁正在学 Transformer 想往视觉方向转的同学、做图像分类任务但觉得 CNN 已经到头了的工程师、以及所有想搞清楚 ViT 内部到底怎么回事的读者。我会先把论文核心段落翻译成中文再逐段拆解里面的设计逻辑最后附上我自己复现时的代码核心片段和踩坑记录保证你能从原理到实践完整走通。1. 先解决翻译这篇论文到底在说什么既然标题带“论文翻译”我先把论文最核心的几个部分的翻译放在前面。我没有逐字逐句全篇翻译那样子太长而且很多铺垫对理解没帮助我挑的是摘要、引言核心段、方法主体、实验结论这几块把论文真正有价值的信息全部覆盖到。1.1 摘要翻译一句话概括 ViT 做了什么原文学术表达比较拗口我翻译得尽量口语化但保留技术准确性虽然在自然语言处理领域Transformer 架构已经成为事实上的标准但在计算机视觉领域Transformer 的应用仍然受到限制。大多数研究中要么保持 CNN 的整体结构不变仅仅用自注意力模块替换其中的某些组件要么用 Transformer 辅助 CNN 但整体上仍然离不开卷积结构。这种依赖是有原因的自注意力是全局操作而图像是像素级别的数据直接对每个像素做自注意力计算复杂度是像素数的平方普通硬件根本扛不住。我们在这篇论文里给出的答案是没必要对像素做注意力把图像切成 16x16 的 patch 网格每个 patch 拉平成向量当作 Transformer 输入序列中的一个 token 就够了。这个做法本质上就是把 NLP 里“词”的概念直接搬到图像上patch 就是图像的“词”。我们把这个模型叫做 Vision TransformerViT。在中等规模的数据集比如 ImageNet上直接训练ViT 的精度比同等规模的 ResNet 要低一些大概差几个点这个结果并不意外因为 Transformer 天生缺乏 CNN 那种内置的归纳偏置比如局部性和平移等变性。但是当训练数据量上到 1400 万张图片ImageNet-21k甚至 3 亿张图片JFT-300M这个级别时情况完全反转ViT 在多个图像识别基准上都超过了当时最先进的 CNN而且训练所需的计算资源还少得多。我们这个模型的设计尽可能贴近原始 Transformer几乎不需要针对视觉任务做特殊改造这意味着 NLP 领域积累的 Transformer 工程优化经验可以直接迁移到视觉任务上来。这段摘要信息量很大。注意最后一句ViT 最大的贡献之一其实是“设计极简”——它没有发明新的模块就是把 BERT 那套东西拿过来改一下输入就行。这降低了整个领域的学习成本也加速了后续研究的爆发。1.2 引言翻译为什么之前没人这么做引言部分我重点翻译关于动机和挑战的段落自注意力架构尤其是 Transformer已经成为 NLP 领域的首选模型。最主流的做法是在大规模文本语料上预训练然后对下游任务做微调。这套“预训练 微调”的范式在 NLP 上极其成功得益于 Transformer 极高的并行性和训练效率。把这套方法搬到视觉领域最大的障碍在于序列长度。一张标准分辨率 224x224 的图片有 50176 个像素如果以像素为 token序列长度就是 5 万。Transformer 的自注意力计算量是序列长度的平方5 万长度的序列意味着 25 亿对像素之间的注意力计算这在现有硬件上完全不可行。之前的研究大多绕开这个问题办法是只在 CNN 的后期特征图上应用注意力或者用局部注意力窗口限制计算范围。这些方法虽然有效但本质上还是“CNN 为主、注意力为辅”并没有摆脱卷积结构。我们的做法是把图片切成固定大小的 patch比如 16x16这样一张 224x224 的图片就变成 196 个 patch序列长度从 5 万降到 196Transformer 直接就能处理。这个设计简单得让人怀疑为什么之前没人试过但实验证明它就是有效。这里有一个值得注意的点为什么之前没人这么做其实之前不是完全没人做过但都很谨慎。2018 年左右有一篇论文做过类似尝试把图像 patch 丢给 Transformer但在中等数据集上效果不如 CNN就没了下文。ViT 的团队坚持做下去并且把“数据规模”这个变量放进来才撬动了整个领域。这告诉我们有时候实验结果不理想可能不是方法本身的问题而是实验条件没到位。1.3 方法部分翻译ViT 的结构详解方法部分是论文最核心的内容我完整翻译了模型结构描述模型结构尽量遵循原始 Transformer 的设计。标准 Transformer 接收一维的 token 嵌入序列作为输入为了处理二维图像我们把图像 x 重塑为一个展平的 2D patch 序列。具体来说输入图像尺寸为 H x W x Cpatch 尺寸为 P x P那么 patch 数量 N HW/P²。每个 patch 被展平成向量再通过一个可训练的线性投影映射到 D 维空间这个投影的输出叫做 patch embedding。为了和 BERT 保持一致我们在序列最前面加了一个可学习的 class token它在 Transformer 编码器输出端的对应状态就作为图像的聚合表示后面接一个分类头。位置编码采用标准的可学习 1D 位置嵌入直接加到 patch embedding 上。之所以用 1D 而不是 2D 位置编码是因为实验发现 2D 位置编码并没有带来明显的精度提升说明模型能够自己学会理解 patch 之间的空间关系。Transformer 编码器由多层组成每层包含标准的多头自注意力MSA和 MLP 块。MLP 由两个全连接层组成中间用 GELU 激活每个模块前面加 LayerNorm模块之间用残差连接。整体结构为LN - MSA - 残差 - LN - MLP - 残差。这里最关键的几个设计决策第一class token 的方式直接借用了 BERT 的 [CLS] token 思想。它在所有 patch embedding 之上额外增加一个可学习的向量经过多层编码器之后这个位置的输出向量聚合了全局信息再接分类头。当然也可以不用 class token直接对所有 patch 的输出做全局平均池化论文里对比过两者效果接近。第二1D 位置编码在直觉上好像会丢失二维空间信息但实验结果说明模型能自行学会这种映射关系。这也解释了 Transformer 强大的表征能力——有些结构上的“常识”其实并非必要。2. ViT 核心原理拆解为什么图片能变成 16x16 个词翻译解决的是“论文说了什么”这一节解决“为什么这样说、底层逻辑是什么”。只有把原理吃透你才能理解为什么 ViT 在大数据上碾压 CNN、为什么小数据上反而打不过。2.1 Patch Embedding切图的本质是一种下采样Patch Embedding 是 ViT 和 CNN 最根本的分水岭。CNN 的核心操作是卷积卷积核在图像上滑动每次只看到局部区域通过层层堆叠不断增大感受野。而 ViT 一步到位直接把整张图切成互不重叠的小块每个小块只经过一次线性变换就变成了 token。拿 16x16 的 patch 尺寸为例一张 224x224 的 RGB 图片会被切成 14x14 的网格总共 196 个 patch。每个 patch 有 16x16x3768 个像素值展平后通过一个全连接层映射到 D 维ViT-Base 中 D768。这个映射其实就是矩阵乘法输入 768 维输出 768 维参数量约 59 万。对比一下ResNet-50 的第一个卷积层参数量约 9400虽然 ViT 第一层参数量大一些但整体网络结构比 CNN 简单得多。“切图”这个操作本质上就是一种下采样——把 50176 个像素点压缩成 196 个 token。信息量肯定有损失但实验证明对于图像分类这种任务保留的细粒度信息完全够用。更有意思的是patch 的多尺度信息其实可以通过调整 patch 尺寸来控制patch 越小序列越长模型越精细但计算量也越大。16x16 是精度和效率的折中。2.2 Position Embedding给 Transformer 补上空间感Transformer 的自注意力机制是“排列不变”的——你把 token 的顺序打乱输出的每个 token 表示理论上不会改变因为注意力权重是按内容相似度算的。但图像的空间结构显然是有意义的一只猫的眼睛在脸上不可能跑到鼻子下面去。所以必须显式地把位置信息注入进去。ViT 的做法是加一个可学习的位置编码矩阵形状是 (N1, D)其中 N196 是 patch 数量1 是 class token 的位置D 是嵌入维度。训练初期这个矩阵是随机初始化的在训练过程中通过反向传播自动学习每个位置应该有什么样的表示。你可能会有疑问为什么不用 NLP 里常用的三角函数位置编码两个原因。第一图像 patch 数量和 NLP 中的 token 数量相比很小196 vs 512 或更多可学习编码的参数量并不大第二论文作者实验过2D 位置编码分别编码行和列在精度上没有显著优势反而增加复杂性所以直接选了最简方案。这里有个细节值得留意位置编码到底编码了什么可视化研究发现学习到的位置编码向量之间的距离呈现有趣的模式——相邻 patch 的位置编码相似度高相距远的 patch 编码差异大。也就是说模型自己“学会”了空间邻近性不需要人为指定。2.3 Transformer Encoder注意力如何在图像上起作用Encoder 部分和 NLP 中的 Transformer 几乎一模一样。每个 Encoder Block 包含两个核心子层多头自注意力MSA和 MLP前置 LayerNorm后接残差连接。自注意力机制在图像上的直观理解是每个 patch 会去看所有其他 patch计算自己与别人的相关性然后根据相关性加权融合全局信息。比如“眼睛”这个 patch 可能会对“嘴巴”patch 有较高的注意力权重因为它们在语义上是相关的。多层叠加之后低层 attention 可能关注颜色纹理中层关注局部形状高层关注物体级别的语义——这种层级语义的涌现不需要人为设计纯粹从数据中学习。多头注意力的“多头”意思是让模型同时从多个子空间去计算关联。ViT-Base 有 12 个头每个头关注不同的关系模式有的头重点关注相邻 patch 的位置关系有的头关注颜色相似的区域有的头关注语义相关的远距离 patch。这种多视角的建模能力让模型能同时捕捉图像中的局部纹理和全局结构。MLP 模块是逐 token 独立的全连接层ViT-Base 中隐藏层维度是 3072是嵌入维度的 4 倍包含两个全连接层中间用 GELU 激活最后加一个 Dropout。这个 MLP 的参数量占整个模型的很大比例实际作用是对每个 token 的表示做非线性变换增强模型的表达能力。2.4 为什么 ViT 在小数据上打不过 CNN大数据上却反超这是理解 ViT 最重要的一个问题。CNN 的核心优势在于内置归纳偏置卷积核的局部连接性假设了邻近像素相关性更高权重共享假设了同一特征在图像不同位置都适用池化操作提供了平移不变性。这些先验知识在小数据集上极其珍贵因为它相当于免费送你一些“经验”不需要从数据中学习。Transformer 没有这些先验。自注意力是全局的它不做任何邻近性假设什么关系都要自己从数据里学。这就意味着模型需要大量数据才能学出“空间邻近”这个常识。在 ImageNet-1k128 万张图上ViT 需要额外正则化手段如数据增强、权重衰减、DropPath才能勉强追平 ResNet预训练数据量一上来这些约束就变成限制——CNN 的归纳偏置反而阻碍了它从更多数据中获益而 Transformer 可以自由地从大数据中学习任何有用的模式。用大白话说CNN 像是自带地图导航的司机小场地不用学就能开但地形一变就受限Transformer 像是只带 GPS 信号但需要自己积累路线的新手司机刚开始开得慢一旦跑够了里程反而能走出一条更优的路径。3. 代码实操从零实现一个 ViT原理清楚之后最好的巩固方式就是亲手写一遍。这一节我会用 PyTorch 从零搭建一个 ViT不借助 torchvision 的高层封装让你真正看到每个组件是什么、怎么组合的。3.1 环境准备建议使用 PyTorch 2.0 以上版本Python 3.9。完整依赖如下pip install torch2.0 torchvision0.15 pip install einops matplotlib tensorboardeinops 是一个张量操作库它的 rearrange 函数可以非常直观地完成 patch 切分和展平操作代码可读性高很多。3.2 Patch Embedding 实现先写最核心的 Patch Embedding。你可能会想切 patch 再拉平再全连接写起来是不是很复杂其实用 PyTorch 的 Conv2d 可以一步搞定原理是一个 kernel_sizestrideP 的卷积等价于把图像切成 P x P 的不重叠 patch然后对每个 patch 做线性投影。import torch import torch.nn as nn from einops import rearrange class PatchEmbed(nn.Module): 将图像切分为 patch 并映射到 embedding 空间 img_size: 输入图像尺寸默认 224 patch_size: patch 尺寸默认 16 in_chans: 输入通道数RGB 图像为 3 embed_dim: embedding 维度ViT-Base 为 768 def __init__(self, img_size224, patch_size16, in_chans3, embed_dim768): super().__init__() self.img_size img_size self.patch_size patch_size self.num_patches (img_size // patch_size) ** 2 self.proj nn.Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): # x: (B, C, H, W) x self.proj(x) # (B, embed_dim, H/P, W/P) x rearrange(x, b c h w - b (h w) c) # (B, num_patches, embed_dim) return xConv2d 的用法要仔细体会kernel_size16stride16卷积核在图上滑动时每次覆盖一个 16x16 的区域且不重叠输出通道数就是 embedding 维度。这和“对每个 patch 做线性映射”在数学上是完全等价的但卷积实现效率高得多因为底层优化做得好。3.3 Transformer Encoder 实现接下来是 Encoder Block包含 LayerNorm、Multi-Head Self-Attention、MLP、残差连接四部分。class MultiHeadSelfAttention(nn.Module): def __init__(self, dim, num_heads12, qkv_biasTrue, attn_drop0., proj_drop0.): super().__init__() self.num_heads num_heads head_dim dim // num_heads self.scale head_dim ** -0.5 self.qkv nn.Linear(dim, dim * 3, biasqkv_bias) self.attn_drop nn.Dropout(attn_drop) self.proj nn.Linear(dim, dim) self.proj_drop nn.Dropout(proj_drop) def forward(self, x): B, N, C x.shape qkv self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads) qkv qkv.permute(2, 0, 3, 1, 4) # (3, B, heads, N, head_dim) q, k, v qkv.unbind(0) attn (q k.transpose(-2, -1)) * self.scale # (B, heads, N, N) attn attn.softmax(dim-1) attn self.attn_drop(attn) x (attn v).transpose(1, 2).reshape(B, N, C) x self.proj(x) x self.proj_drop(x) return x class MLP(nn.Module): def __init__(self, in_dim, hidden_dim, act_layernn.GELU, drop0.): super().__init__() self.fc1 nn.Linear(in_dim, hidden_dim) self.act act_layer() self.fc2 nn.Linear(hidden_dim, in_dim) self.drop nn.Dropout(drop) def forward(self, x): x self.fc1(x) x self.act(x) x self.drop(x) x self.fc2(x) x self.drop(x) return x class TransformerBlock(nn.Module): def __init__(self, dim, num_heads, mlp_ratio4., drop0., attn_drop0.): super().__init__() self.norm1 nn.LayerNorm(dim) self.attn MultiHeadSelfAttention(dim, num_heads, attn_dropattn_drop, proj_dropdrop) self.norm2 nn.LayerNorm(dim) self.mlp MLP(dim, int(dim * mlp_ratio), dropdrop) def forward(self, x): x x self.attn(self.norm1(x)) x x self.mlp(self.norm2(x)) return x这里有个很容易写错的地方标准的 Pre-LN 结构是“先 LayerNorm再 Attention再接残差”和 Post-LN先计算再归一化顺序不同。ViT 用的是 Pre-LN这样做的好处是训练更稳定梯度传播更顺畅即使层数很深也不容易梯度消失。注意力计算公式 Q K^T 之后要除以 sqrt(head_dim)这是为了缩放点积防止注意力分数过大导致 softmax 之后梯度消失。head_dim 768 / 12 64所以 scale 1/8。3.4 完整 ViT 模型组装把 Patch Embedding、Class Token、位置编码和若干个 Transformer Block 组装起来class ViT(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, num_classes1000, embed_dim768, depth12, num_heads12, mlp_ratio4., drop0., attn_drop0.): super().__init__() self.patch_embed PatchEmbed(img_size, patch_size, in_chans, embed_dim) num_patches self.patch_embed.num_patches 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(drop) self.blocks nn.Sequential(*[ TransformerBlock(embed_dim, num_heads, mlp_ratio, drop, attn_drop) 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, std0.02) nn.init.trunc_normal_(self.cls_token, std0.02) self.apply(self._init_weights) def _init_weights(self, m): if isinstance(m, nn.Linear): nn.init.trunc_normal_(m.weight, std0.02) if m.bias is not None: nn.init.zeros_(m.bias) elif isinstance(m, nn.LayerNorm): nn.init.ones_(m.weight) nn.init.zeros_(m.bias) def forward(self, x): B x.shape[0] x self.patch_embed(x) # (B, N, embed_dim) cls_tokens self.cls_token.expand(B, -1, -1) x torch.cat([cls_tokens, x], dim1) # (B, N1, embed_dim) x x self.pos_embed x self.pos_drop(x) x self.blocks(x) x self.norm(x) # 取 class token 对应的输出 cls_output x[:, 0] return self.head(cls_output) # 创建一个 ViT-Base/16 实例 model ViT(img_size224, patch_size16, num_classes1000, embed_dim768, depth12, num_heads12) # 测试前向传播 dummy torch.randn(2, 3, 224, 224) output model(dummy) print(output.shape) # torch.Size([2, 1000])Position Embedding 的实际形状是 (1, 197, 768)其中 197 1class token 196patch。维度 0 的 batch 维用 expand 复制到每个样本注意这里不能用 repeatexpand 是共享内存的省资源。3.5 训练要点与代码级调优模型写完了训练时有一些经验值得分享优化器选择ViT 适合 AdamW 而不是纯 Adam。AdamW 把权重衰减和梯度更新解耦对 Transformer 这类大模型更友好。学习率建议 1e-3 到 3e-3配合线性预热warmup和余弦退火。我在 ImageNet-1k 上复现时用了 5000 步 warmup初始学习率 1e-6 逐步升到 1e-3效果比固定学习率好很多。数据增强策略ImageNet-1k 规模的数据集对 ViT 来说是“小数据”必须上强增强。我实测下来RandAugment Mixup CutMix Random Erasing 这几招组合起来能让 ViT-Base 在 ImageNet-1k 上的 top-1 提升 2~3 个点。没有这些增强ViT 在小数据上就是欠拟合加过拟合的双重灾难。from timm.data import create_transform transform create_transform( input_size224, is_trainingTrue, auto_augmentrand-m9-mstd0.5-inc1, mixup_alpha0.8, cutmix_alpha1.0, re_prob0.25, )Batch Size 的影响ViT 对 batch size 比较敏感。在单卡 2080Ti11G上batch size 只能到 64 左右我用的是梯度累积模拟更大的 batch。实测 batch size 从 256 升到 1024相同 epoch 下 top-1 能提升约 1%。所以条件允许的话尽量用大 batch 配合大学习率。学习率缩放规则线性缩放法则Linear Scaling Rule在这里同样适用——batch size 翻倍学习率也翻倍。我从 batch 512 换到 1024 时学习率从 1e-3 调整到 2e-3训练曲线才恢复正常。4. 常见问题与排查技巧实录从我自己的复现经历和帮助过的同学反馈来看训练 ViT 时遇到的问题主要集中在以下几个方面我整理成了一份问题速查表4.1 问题速查表问题现象可能原因解决方案训练 loss 不下降学习率过大/过小或位置编码未正确初始化检查 warmup 策略学习率从 0 线性升到 1e-3确认 pos_embed 使用了 trunc_normal_ 初始化小数据集上严重过拟合Transformer 缺少归纳偏置需要更多数据或正则增加数据增强Mixup、CutMix、RandAugment增大 DropPath 概率到 0.1~0.3添加随机深度显存不足 OOM自注意力矩阵占用显存随序列长度平方增长降低 batch size使用梯度累积用 torch.utils.checkpoint 做激活检查点训练速度极慢未使用混合精度或者 PyTorch 版本过低用 AMPtorch.cuda.amp混合精度训练速度提升 2~3 倍微调时精度反而下降学习率设置不当或未做位置编码插值适配微调学习率建议 3e-5 到 1e-4输入分辨率改变时需要对 pos_embed 做插值注意力训练不稳定LayerNorm 位置错误或初始化方式不当确认使用 Pre-LN 结构LN 在 Attention 之前Linear 层用 trunc_normal_ 初始化4.2 高频问题深入解析问题一为什么小数据集上 ViT 就是训不过 ResNet这是新手最容易困惑的地方其实原理在我前面的分析中已经提到了。Transformer 的归纳偏置弱等于“从零开始学所有规则”CNN 天生带局部性和平移不变性。128 万张图的 ImageNet-1k 对 Transformer 来说不够学但对 CNN 来说已经“喂饱了”。解决办法除了加大数据还可以用蒸馏或者混合架构如 DeiT、T2T-ViT或者直接用预训练好的权重做微调。问题二为什么我改输入分辨率后模型崩溃了ViT-Base 预训练时用的是 224x224如果你在推理时输入 384x384patch 数量就从 196 变成 576但位置编码矩阵的形状还是 (1, 197, 768)维度不匹配程序直接报错。解决办法是“位置编码插值”def interpolate_pos_embed(pos_embed, new_num_patches): 将预训练位置编码插值到新的 token 数量 pos_embed pos_embed.unsqueeze(0) # (1, 197, 768) pos_embed pos_embed.permute(0, 2, 1) # (1, 768, 197) pos_embed nn.functional.interpolate( pos_embed, sizenew_num_patches, modebilinear, align_cornersFalse) pos_embed pos_embed.permute(0, 2, 1) return pos_embed注意操作顺序先去掉 class token 的位置第 0 个插值完成后再拼回去。直接对所有位置做插值会把 class token 的数据也混进去导致信息污染。问题三ViT 训练非常慢有没有提速方法有而且很有效。第一是混合精度FP16 在 Ampere 架构显卡上能带来 2 倍以上的加速第二是 Flash AttentionPyTorch 2.0 自带的torch.nn.functional.scaled_dot_product_attention内部已经实现了 Flash Attention 相关的优化能大幅减少显存占用并提高注意力计算速度第三是 use_better_transformermodel.to(memory_formattorch.channels_last)在部分硬件上有额外加速。我在 3090 上实测这三个方法合计能省一半以上训练时间。问题四Dropout 和 DropPath 怎么设置ViT 中有两种常用的随机正则手段。Dropout 作用于 embedding 和 MLP 输出BERT 时代的标准做法DropPath 针对残差分支训练时随机丢弃某些层的小数路径相当于深度上的正则。官方 timm 实现里ViT-Base/16 的 drop_path_rate 一般设为 0.1DeiT 系列设为 0.1~0.3。小数据集建议调大大数据集反而要调小或不用因为大数据本身已有足够正则效果。5. 延伸思考从 ViT 到 Deformable DETR 的视觉 Transformer 进化写完一个 ViT 之后如果你还想继续深入我建议顺着一条线索走ViT 解决了“图像分类中的全局建模”那检测、分割这类密集预测任务怎么办这就延伸到了 Deformable DETR也是标题里提到的热词。Deformable DETR可变形 DETR的核心思想是解决 DETR 收敛慢的问题。原始 DETR 用 Transformer 做端到端目标检测但注意力模块在初始化时均匀关注整张图的特征点导致需要很长的训练周期才能让注意力聚焦到物体上。Deformable DETR 的做法是借鉴可变形卷积的思路只让每个 query 关注一组稀疏的采样点而不是全图所有位置。这个方法背后的逻辑和 ViT 有异曲同工之处都是想办法降低 Transformer 在视觉任务上的计算复杂度。ViT 用 patch 降低序列长度Deformable DETR 用稀疏采样降低注意力矩阵大小。这里有个值得注意的设计细节Deformable DETR 里的“可变形”不是传统可变形卷积那种学习 offset 的做法而是通过一个轻量网络预测每个 query 在特征图上采样的若干参考点然后只在参考点周围做注意力。参考点数量通常是 4 或 8远小于特征图上的 10000 位置计算量大幅下降收敛速度比 DETR 快了近 10 倍。如果你对 Deformable DETR 的实现感兴趣核心在于两点多尺度特征图上做可变形注意力、以及用二分图匹配做检测框匹配。代码实现上可变形注意力的 CUDA 算子比较难手写建议直接使用官方或者 mmdetection 中的实现。从 ViT 到 Deformable DETR能看到一条清晰的进化线索Transformer 进入视觉领域后核心问题始终围绕着“如何控制计算复杂度”和“如何融入视觉先验”这两个维度展开。ViT 通过 patch 化解决前者Deformable DETR 通过稀疏注意力解决前者并顺便解决了收敛问题。理解这条线索对你阅读后续的 Swin Transformer、DINO、SAM 等模型会非常有帮助。6. 复现过程中的踩坑记录与心得体会最后分享几个我实践中印象深刻的坑希望你能绕开。第一个坑和位置编码的初始化有关。我最早实现时位置编码用了简单的高斯随机初始化mean0, std1结果训练前期 loss 很难下降仔细排查才发现是位置编码初始方差过大导致 attention score 在初期的信噪比太低。后来改成trunc_normal_(std0.02)问题立刻消失。timm 里的实现也是这么做的初始化标准差和 Linear 层保持同一量级这个细节千万别忽略。第二个坑是关于冻结预训练权重的微调策略。我在做迁移学习时一开始把整个网络都冻结只训最后的分类头结果精度还不如从零训练。后来发现对于 Transformer 架构分类头只占模型很小一部分参数量冻结主干等于浪费了大部分可调参数。正确的做法是如果数据量很小可以只解冻最后几层如果数据量中等全量微调但用更低的学习率1e-4 或更低配合早停。第三个坑更加隐蔽图像归一化。ViT 预训练时用的是 ImageNet 的 mean[0.485, 0.456, 0.406] 和 std[0.229, 0.224, 0.225]如果你在推理时忘了做同样的归一化精度会大幅下降甚至掉十几个点。这个问题我在工程部署时踩过一次排查了很久最后发现是预处理环节漏了一步。第四个坑和第 2D 位置编码有关。我之前一直觉得 2D 位置编码在理论上应该更好毕竟它提供了行列信息就自己改了实现。结果在相同实验条件下2D 编码和 1D 编码精度几乎没有差别甚至收敛略慢。后来我明白了Transformer 的注意力机制本身已经足够灵活空间信息可以通过特征相似度隐式地学习到人为强加的结构反而可能限制它的自由度。这也是论文作者选择 1D 编码的原因——不是为了简化而是实验证明默认方案已经够好。根据我个人经验复现 ViT 的时候不要一上来就追求跟论文完全一致的效果。先在小数据集CIFAR-10 或 CIFAR-100上调通流程确认模型和训练代码没有 bug再上 ImageNet 级别的数据和预训练。我在 CIFAR-10 上花了 2 天时间摸索出正确的增强策略和学习率然后迁移到 ImageNet 只跑了 3 天就到了可接受精度约 79% top-1整个过程比一开始硬啃大数据集高效得多。最后再分享一个小技巧如果你想在视觉任务上快速试用 ViT建议直接用 timm 库中的实现。timm 团队的实现经过了大面积验证包含了各种优化细节比如 head 初始化为零、DropPath 的缩放逻辑比自己从零写的版本稳定很多。但我会建议你还是至少手动实现一次因为只有自己写过一遍你才能真正理解 ViT 的设计哲学——也才能在它出问题时知道该从哪里下手排查。
返回列表