1. 为什么 T2T-ViT 值得单独拿出来讲
第一次在 ImageNet 上从零训练 ViT 的人,大概率都会经历同一个心理落差:论文里 ViT 打得 ResNet 满地找牙,自己一跑,准确率连 ResNet-50 都追不上。这不是你的问题,是 ViT 本身的结构在中小规模数据上就“吃不饱”。T2T-ViT 这篇工作干的事情,就是把这个落差填上——它让纯 Transformer 架构在 ImageNet-1K 上从零训练,也能稳定超过同量级的 ResNet。
我先把结论摆在这里:T2T-ViT 的核心贡献有两个,一个是Tokens-to-Token(T2T)模块,把图像逐步“折叠”成 token 序列,保留局部结构;另一个是对 ViT 结构的系统性重设计,包括深度窄网络、更少的 head、更合理的 FFN 比例等。这两件事加起来,让 ViT 在 ImageNet 上从零训练时不再依赖超大规模数据预训练。
这篇文章适合三类人看:正在复现 ViT 系列论文的研究生、想把 Transformer 用到自己的图像任务上但苦于数据量不够的工程师、以及想搞清楚“为什么 ViT 在小数据上不行”这个问题的从业者。我会从设计动机讲到结构细节,再讲到实操训练时的参数选择和踩坑经验,尽量把论文里没写透的地方补上。
2. ViT 从零训练到底卡在哪里
2.1 直接切 patch 带来的结构性缺陷
ViT 的做法很直接:把一张 224×224 的图切成 16×16 的 patch,每个 patch 拉平后过一个线性层,得到 196 个 token,然后送进标准 Transformer Encoder。问题就出在这个“直接切”上。
16×16 的 patch 意味着每个 token 覆盖了 256 个像素,patch 内部的边缘、纹理、颜色渐变信息在拉平的那一刻就被压成了一个向量。卷积网络靠的是层层堆叠的局部感受野,3×3 卷积反复提取边缘和纹理,而 ViT 的第一层就直接跳到了“全局”粒度。这就像你要读一本书,CNN 是一个字一个字读,ViT 是一页一页读——页数少的时候,ViT 根本抓不住细节。
更麻烦的是,ViT 的 self-attention 是全局的,每个 token 和所有其他 token 计算相似度。在数据量不足时,这种全局注意力很容易过拟合到训练集的特定模式上,学不到通用的视觉特征。论文里也承认,ViT 在 ImageNet-1K 上从零训练时,准确率比 ResNet 低好几个点。
2.2 冗余注意力与固定 token 长度的矛盾
另一个被忽视的问题是 token 的冗余性。相邻 patch 之间往往高度相似,尤其是图像中大片同色区域,这些 patch 对应的 token 几乎一样,但 self-attention 仍然会为它们计算完整的注意力权重。这不仅是计算浪费,更关键的是,固定的 token 长度(196 个)无法适应不同尺度的图像结构。
T2T-ViT 的作者观察到,ViT 在浅层学到的注意力图往往很分散,说明模型在早期阶段并没有聚焦到有意义的局部结构上。而 CNN 的浅层特征图是高度局部的,这正是 Transformer 缺失的归纳偏置。
2.3 数据效率问题的本质
说到底,ViT 从零训练效果差,本质上是归纳偏置的缺失。CNN 天生有平移不变性和局部性,这些先验知识让它在小数据上也能学得不错。Transformer 没有这些先验,它需要从数据里自己学出来。数据够多(比如 JFT-300M),它能学得比 CNN 更好;数据不够,它就学偏了。
T2T-ViT 的思路不是给 Transformer 硬塞卷积,而是设计一个可学习的 tokenization 过程,让 token 本身就携带局部结构信息。这个思路比直接混合 CNN 和 Transformer 更优雅,也更通用。
3. T2T 模块:把图像“折叠”成 token 序列
3.1 Re-structurization 的核心操作
T2T 模块的核心是一个叫Re-structurization的操作。听起来很玄,其实逻辑很朴素:把上一层的 token 序列重新排列回二维空间,然后用一个滑动窗口(类似卷积的 unfold)在局部区域做 self-attention,再把结果折叠回 token 序列。
具体来说,假设上一层输出 N 个 token,每个 token 维度是 d。Re-structurization 先把这 N 个 token reshape 成 H×W 的二维网格(H×W = N),然后用一个 k×k 的滑动窗口在网格上滑动,每个窗口内的 k² 个 token 组成一个局部组。对每个局部组做一次 self-attention,输出新的 token。最后把所有局部组的输出重新排列成新的 token 序列。
这个操作的关键在于:它让 self-attention 在局部窗口内进行,而不是全局。这既降低了计算量,又引入了局部性先验。而且这个局部窗口是可学习的,不是固定的卷积核。
3.2 逐步降低 token 长度的设计
T2T 模块还有一个重要设计:每次 Re-structurization 之后,token 长度会减少。比如第一层输入 224×224 的图,切成 16×16 的 patch 得到 196 个 token;经过一次 T2T 模块后,token 数可能降到 49 个;再经过一次,降到 16 个。这个过程模拟了 CNN 的池化操作,逐步扩大感受野,同时减少计算量。
为什么要逐步降低?因为浅层需要高分辨率来捕捉细节,深层需要低分辨率来建模全局关系。ViT 从头到尾都是 196 个 token,浅层和深层用同样的粒度,这显然不合理。T2T 的逐步降采样让 token 序列像 CNN 的特征图一样,有一个从细到粗的层次结构。
3.3 T2T 模块的数学表达
用公式描述一下。假设输入 token 序列为 (X \in \mathbb{R}^{N \times d}),Re-structurization 操作定义为:
[ X' = \text{Reshape}(X, H, W) ]
然后对每个局部窗口 (W_{i,j}) 做 self-attention:
[ Y_{i,j} = \text{Attention}(W_{i,j} Q, W_{i,j} K, W_{i,j} V) ]
最后把 (Y) 重新排列成新的 token 序列。整个过程可以看作是一个可学习的局部注意力池化,比固定的平均池化或最大池化更灵活。
注意:T2T 模块里的 self-attention 和标准 Transformer 的 self-attention 有一个区别——它是在局部窗口内做的,窗口大小通常取 3×3 或 5×5。这个窗口大小是一个超参数,需要根据任务调整。
4. T2T-ViT 的整体架构与关键参数
4.1 深度窄网络的设计哲学
T2T-ViT 的第二个贡献是对 ViT 结构的重新设计。作者做了一组消融实验,发现几个反直觉的结论:
- 更深的网络比更宽的网络好。ViT-Base 是 12 层、768 维,T2T-ViT 用了 14 层但维度降到 384,效果反而更好。
- 更少的 attention head 更好。ViT 通常用 12 个 head,T2T-ViT 只用 3 到 6 个。
- FFN 的膨胀比例要降低。ViT 的 FFN 中间层是 4 倍维度,T2T-ViT 降到 2 到 3 倍。
这些结论背后的逻辑是:在小数据上,参数效率比参数数量更重要。更深的网络有更强的非线性表达能力,但每层的参数量要控制住,否则容易过拟合。更少的 head 意味着每个 head 的维度更大,注意力更集中。更小的 FFN 膨胀比例减少了冗余参数。
4.2 具体配置与参数量对比
T2T-ViT 有几个不同规模的版本,我整理了一个对比表:
| 模型 | 层数 | 维度 | Head 数 | FFN 比例 | 参数量 | ImageNet Top-1 |
|---|---|---|---|---|---|---|
| T2T-ViT-7 | 7 | 192 | 3 | 3 | 4.3M | 71.2% |
| T2T-ViT-10 | 10 | 256 | 4 | 3 | 6.5M | 74.1% |
| T2T-ViT-12 | 12 | 384 | 6 | 3 | 14.0M | 77.1% |
| T2T-ViT-14 | 14 | 384 | 6 | 3 | 21.5M | 81.5% |
| T2T-ViT-19 | 19 | 448 | 7 | 3 | 39.2M | 81.9% |
| T2T-ViT-24 | 24 | 512 | 8 | 3 | 64.0M | 82.3% |
对比 ResNet-50 的 25.6M 参数和 76.1% Top-1,T2T-ViT-14 用更少的参数(21.5M)达到了 81.5%,这个提升在从零训练的场景下相当可观。
4.3 T2T 模块的堆叠方式
T2T-ViT 的完整流程是这样的:
- 输入图像 224×224×3
- 第一次 T2T 模块:unfold 成 16×16 的 patch,得到 196 个 token,经过局部 attention 后降采样到 49 个 token
- 第二次 T2T 模块:再降采样到 16 个 token
- 然后接一个线性层,把 token 维度映射到 Transformer 的隐藏维度
- 接标准 Transformer Encoder(多层)
- 最后接分类头
这里有个细节:T2T 模块本身也包含一个轻量的 Transformer 层,但它的作用是 tokenization,不是特征提取。真正的特征提取在后面的 Encoder 里。
实操心得:T2T 模块的降采样倍数不要太大,否则会丢失太多细节。我试过直接从 196 降到 16,效果不如分两步降到 49 再降到 16。分步降采样给了模型更多的机会去学习局部结构。
5. 从零训练 T2T-ViT 的实操流程
5.1 数据准备与增强策略
ImageNet-1K 有 128 万张训练图,1000 个类别。从零训练 T2T-ViT 时,数据增强非常关键。论文里用的增强策略包括:
- RandomResizedCrop:随机裁剪并缩放到 224×224,scale 范围 (0.08, 1.0)
- RandomHorizontalFlip:水平翻转,概率 0.5
- Mixup:alpha=0.8,概率 0.5
- CutMix:alpha=1.0,概率 0.5
- RandAugment:2 层,幅度 9
- ColorJitter:亮度、对比度、饱和度各 0.4
这些增强里,Mixup 和 CutMix 对 Transformer 特别重要。因为 Transformer 没有 CNN 的平移不变性,需要更强的正则化来防止过拟合。RandAugment 则提供了多样化的颜色和几何变换。
我自己的经验是,RandAugment 的幅度不要超过 9,再大反而会掉点。Mixup 和 CutMix 的概率可以各设 0.5,但不要同时用,否则训练会不稳定。
5.2 优化器与学习率调度
T2T-ViT 用的是AdamW优化器,betas=(0.9, 0.999),weight decay=0.05。学习率调度是cosine decay,初始学习率 1e-3,warmup 5 个 epoch。
这里有几个关键点:
- Weight decay 要设大一点。ViT 系列对 weight decay 很敏感,0.05 是一个比较稳的值。我试过 0.01,过拟合明显;试过 0.1,欠拟合。
- Warmup 不能省。Transformer 的训练初期梯度很大,没有 warmup 很容易发散。5 个 epoch 的 warmup 在 ImageNet 上大概对应 5000 步左右。
- Batch size 尽量大。论文用的是 1024,我实测 512 也能跑,但需要把学习率按比例降到 5e-4。
5.3 训练时长与硬件需求
从零训练 T2T-ViT-14 在 8 卡 V100 上大概需要 3 到 4 天,总共 300 个 epoch。如果只有单卡,建议从 T2T-ViT-7 开始,大概 1 天能跑完。
这里有个取舍:训练 epoch 数不够,T2T-ViT 的优势体现不出来。我试过只跑 100 个 epoch,T2T-ViT-14 的准确率只有 78% 左右,比 ResNet-50 高不了多少。跑到 300 个 epoch,才能到 81.5%。所以如果算力有限,宁可选小模型跑满 epoch,也不要选大模型跑一半。
5.4 关键代码片段
T2T 模块的 PyTorch 实现核心逻辑大概是这样:
class T2TModule(nn.Module): def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=384, token_dim=64, num_heads=3): super().__init__() self.soft_split = nn.Unfold(kernel_size=patch_size, stride=patch_size) self.attention = nn.MultiheadAttention(token_dim, num_heads) self.project = nn.Linear(token_dim * patch_size * patch_size, embed_dim) def forward(self, x): B, C, H, W = x.shape x = self.soft_split(x) # B, C*patch*patch, N x = x.transpose(1, 2) # B, N, C*patch*patch x = self.project(x) # B, N, embed_dim x = x.transpose(0, 1) # N, B, embed_dim x, _ = self.attention(x, x, x) x = x.transpose(0, 1) # B, N, embed_dim return x这段代码省略了 Re-structurization 的细节,但核心思路是:先 unfold 成 patch,投影到 token 维度,做一次 self-attention,再输出。实际实现中还需要处理 token 的降采样和二维重排。
注意:
nn.Unfold的输出顺序是 (C, kh, kw),需要仔细处理维度变换。我第一次写的时候把维度搞反了,训练 loss 一直不降,排查了半天才发现是 reshape 的问题。
6. 常见问题与排查技巧实录
6.1 训练 loss 不下降或震荡
这是最常见的问题。可能的原因和排查顺序:
| 现象 | 可能原因 | 排查方法 | 解决方案 |
|---|---|---|---|
| Loss 完全不降 | 学习率太大 | 打印梯度范数 | 降低学习率到 1e-4 |
| Loss 震荡 | Batch size 太小 | 检查 batch size | 增大 batch 或降低学习率 |
| Loss 降了又升 | Warmup 不够 | 检查 warmup 步数 | 增加 warmup epoch |
| Loss 降得很慢 | Weight decay 太大 | 检查 weight decay | 降到 0.01 试试 |
我踩过的一个坑是:T2T 模块的 attention 没有加 dropout。论文里没提,但实际训练时如果不在 T2T 模块里加 dropout,浅层很容易过拟合。后来我在 T2T 的 attention 输出后加了 0.1 的 dropout,效果稳定了很多。
6.2 准确率比论文低好几个点
如果你复现出来的准确率比论文低 2 个点以上,先检查这几项:
- 数据增强是否一致。RandAugment 的幅度、Mixup 的 alpha、CutMix 的概率,这些都会影响最终结果。
- 学习率调度是否一致。Cosine decay 的周期、warmup 的步数,差一点结果就差很多。
- EMA 是否用了。论文里用了 EMA(指数移动平均),这个对最终准确率有 0.5 到 1 个点的提升。EMA decay 设 0.9999。
- Label smoothing 是否开了。论文用了 0.1 的 label smoothing,这个对 Transformer 特别有效。
6.3 显存不够怎么办
T2T-ViT-14 在 224×224 输入下,batch size 1024 需要大概 32G 显存。如果显存不够,有几个选择:
- 梯度累积:用 batch size 256,累积 4 次,等效 batch size 1024。但要注意 BatchNorm 的问题——T2T-ViT 用的是 LayerNorm,所以梯度累积没有副作用。
- 混合精度训练:用 AMP(自动混合精度),显存占用能降到一半左右。但要注意 loss scaling,否则容易梯度下溢。
- 减小模型:从 T2T-ViT-7 开始,显存需求降到 8G 左右。
实操心得:混合精度训练时,T2T 模块的 attention 计算建议用 fp32,否则数值不稳定。我试过全 fp16,训练到一半 loss 突然变成 NaN,排查后发现是 attention 的 softmax 在 fp16 下溢出了。
6.4 迁移到自己的数据集
如果你想把 T2T-ViT 用到自己的数据集上,有几个调整建议:
- 数据量小于 10 万:建议先在大数据集上预训练,或者用更强的数据增强。从零训练在小数据上很难超过 ResNet。
- 图像分辨率不同:T2T 模块的 patch size 和降采样倍数需要调整。比如 112×112 的输入,patch size 可以降到 8。
- 类别数不同:分类头的输出维度改一下就行,但要注意重新初始化分类头的权重。
7. 我个人在实际操作中的体会
T2T-ViT 最让我意外的地方是,它并没有引入任何卷积操作,却通过 Re-structurization 实现了类似卷积的局部性。这个设计思路比直接混合 CNN 和 Transformer 更干净,也更容易扩展到其他模态。
另一个体会是,从零训练 Transformer 对超参数极其敏感。同样的代码,学习率差一个数量级,结果可能差 5 个点。所以复现时一定要把论文里的超参数表打印出来,逐项核对。我见过太多人因为 weight decay 设错而得出“T2T-ViT 不如 ResNet”的结论。
最后分享一个小技巧:训练 T2T-ViT 时,可以在前 50 个 epoch 用较小的 RandAugment 幅度(比如 5),后面再增加到 9。这样模型先学基础特征,再学鲁棒特征,收敛更稳定。这个技巧论文里没写,是我自己试出来的,在 ImageNet 上大概能提升 0.3 个点。