timm
ViT_Fast / docs /research_story.md
1999xia's picture
Upload folder using huggingface_hub
54ee1eb verified
|
Raw
History Blame Contribute Delete
11.5 kB

让 ViT 跑得更快:我们试了三种方案,最终发现最简单的最有效


问题

ViT-B/16 处理一张 224×224 的图片时,会切成 196 个 patch,全部送入 12 层 transformer。这是 33.85 GFLOPs 的计算量。

但直觉告诉我们:不是所有 patch 都同等重要。背景、天空、重复纹理 —— 这些 patch 对分类贡献很小。如果能在进入 transformer 之前过滤掉不重要的 patch,就能省下大量计算。

于是我们开始了探索。方向很明确,但路走了三条。


基线:全量 ViT-B/16 微调

所有对比的基准是 ViT-B/16(85.8M 参数)在 ImageNet-21K 上预训练后,在目标数据集上全量微调

训练配置:

  • 优化器:AdamW,lr=3e-5,weight_decay=0.05
  • 调度器:CosineAnnealing,100 epochs(CIFAR-100)/ 30 epochs(Food-101)
  • 数据增强:RandomResizedCrop + RandomHorizontalFlip
  • 标签平滑:0.1
  • 有效 batch size:128(batch=32 × accum=4)
数据集 类别 原生分辨率 Test Acc 适合降采样实验
Oxford Pets 37 ~200-1000px(高) 93.81% ✅ 主数据集
Food-101 101 ~300-512px(高) 91.37%
DTD 47 ~300-400px(高) 80.85% ✅ 纹理敏感
CIFAR-100 100 32×32(极低) 91.69% ❌ 不适合降采样

⚠️ 关于 CIFAR-100:原图仅 32×32,训练时放大到 224×224 才喂给 ViT。本身就没多少细节,降采样测试会给出"几乎不掉精度"的假象。CIFAR-100 适合测试 patch 级别的选择策略(如 Router),不适合测试降采样/分辨率相关的实验。所有降采样实验应以 Oxford Pets、Food-101 等高分辨率数据集为准。


网络架构

在进入具体方案前,先理解模型的结构。

基线:全量 ViT-B/16

输入图片 (224×224)
  ↓
Patch Embedding (Conv2d, 16×16, stride=16) → 196 patches, 每个 768维
  ↓
Position Embedding + CLS Token
  ↓
12 层 Transformer Block(每层 = Attention + MLP + Residual)
  ↓
LayerNorm → 分类头 → 输出

Patch Selection 架构(加了 Router)

输入图片 (224×224)
  ↓
Patch Embedding → 196 patches (768维)
  ↓
Position Embedding(位置编码)
  ↓
┌─ Router ─────────────────────────────┐
│ Linear(768→384) → LayerNorm → GELU   │  → 对每个 patch 输出 1 个分数
│ Linear(384→1)                        │
│ 参数量: 768×384 + 384×1 = 295,296     │
└──────────────────────────────────────┘
  ↓ 196 个分数
Differentiable Top-K 选择(Sigmoid STE)
  ↓ 保留 K 个最重要的 patch
concat(CLS Token, 选中 patches) → (K+1, 768)
  ↓
前 2 层 Transformer Block(轻量编码器)
  ↓ 保存中间特征 → 送入 MAE Decoder(训练时)
后 10 层 Transformer Block(主干网络)
  ↓
LayerNorm → 分类头 → 输出

**Router 的参数量(296K)仅为 ViT-B/16 整体(85.8M)的 0.34%**,极小。

Differentiable Top-K 原理

选择"选 K 个 patch"是离散操作,没法直接算梯度。用 STE 技巧绕过:

scores = router(patch_features)        # 196 个分数
threshold = scores.topk(K)[-1:]        # 第 K 大的分数作为阈值
soft_mask = sigmoid(scores - threshold) # 连续的可微分掩码
hard_mask = (scores > threshold).float() # 离散的选/不选
# STE:前向用 hard_mask,反向梯度直穿 soft_mask
mask = hard_mask.detach() + soft_mask - soft_mask.detach()

第一条路:Gumbel Selection —— 学不动

思路: 在 Router 打分后,用 Gumbel-Softmax + STE 做可微分的"选/不选",只把选中的 patch 送进 backbone。

结果(保留 50%):

数据集 Baseline Gumbel Selection 下降
CIFAR-100 91.69% 87.18% -4.51%
Oxford Pets 93.81% 87.00% -6.81%
Food-101 91.37% 85.77% -5.60%

问题在哪? Gumbel 噪声太大了。Router 想表达"这个 patch 重要,score=0.8",但噪声一加,变成了随机数。梯度信号被淹没,Router 学了一个训练周期,最后还是没学会。

教训: 用噪声做离散化选择,噪声强度控制不好就白学。


第二条路:MAE + 蒸馏 Router —— 有效,但复杂

思路: 去掉 Gumbel,改用 Sigmoid STE(梯度直通),让梯度直接穿过"选/不选"的门槛。同时加一个 MAE 重建任务:Encoder 处理后,用小 Decoder 重建被丢弃的 patch,用 MSE Loss 告诉 Router"你丢的东西能不能猜回来"。

这相当于给 Router 配了两个教练:

  • 分类 Loss 告诉它:选出来的 patch 要能分类
  • 重建 Loss 告诉它:丢掉的 patch 不能是关键信息

再进一步: 能不能先让 Router 看看"老师"怎么做?用预训练 ViT 的 CLS 注意力分数预训练 Router,给它一个更好的起点。

主数据集 Oxford Pets(保留 50%):

方法 Test Acc vs Baseline
Baseline(全量) 91.91%
MAE + 从头训练 Router 88.80% -3.11%
MAE + 预热 Router 85.85% -6.06%

跨数据集验证:

数据集 Baseline MAE + 预热 Router(50%) 下降
CIFAR-100 91.69% 89.07% -2.62%
CIFAR-100(75%保留) 91.69% 91.10% -0.59%
Food-101 91.37% 89.52% -1.85%
Oxford Pets 91.91% 85.85% -6.06%

关键发现: 蒸馏效果严重依赖数据集大小。Oxford Pets 仅 3.6K 训练图片,蒸馏信号极弱(注意力重叠率仅 51.16%≈随机),预热反而有害。CIFAR-100 和 Food-101 上蒸馏有效(-1.85%~-2.62%)。

但有个细节耐人寻味: 蒸馏 Router 的注意力重叠率最高仅 59.95%(Food-101),随机基线 50%。几乎没学到老师。

那 MAE 是怎么学到东西的?答案可能是:MAE 重建任务本身就是够好的训练信号,蒸馏的贡献其实很小。Router 在 MAE 训练中自己学会了"什么 patch 能丢",教师注意力只是给了它一个微不足道的热身。

教训: 学习型方案能工作,但训练流程复杂(两阶段),核心收益来自 MAE 而非蒸馏,且小数据集上蒸馏可能有害。


第三条路:降采样训练 —— 简单到让人怀疑

这时我们退了一步,问了一个更基本的问题:

"如果我只是把输入图片缩小一点,会怎么样?"

原图 224×224 → Resize(168×168) → Patch Embed → 100 patches → ViT → 分类

不需要 Router、不需要蒸馏、不需要 MAE。就一行 Resize。

但要小心:CIFAR-100 是个陷阱

最初在 CIFAR-100 上测试,结果漂亮得可疑:

CIFAR-100 原图只有 32×32 → 放大到 224 → 缩回 168 → 几乎没有信息损失

CIFAR-100 的"不掉精度"是假象。换到真正的高分辨率数据集上测试:

真实数据集上的结果

数据集 特点 Baseline 168×168 下降
CIFAR-100 粗粒度,原生 32×32 91.69% 91.56% -0.13% ❌ 假象
Food-101 细粒度菜品 91.37% 89.87% -1.50%
Oxford Pets 猫狗品种,花纹重要 93.81% 90.65% -3.16%
DTD 纯纹理分类 80.85% 74.20% -6.65%

精度下降幅度与数据集的信息依赖类型高度相关:

  • 粗粒度物体(CIFAR-100):降采样几乎没影响——因为只需要轮廓
  • 细粒度物体(Food-101, Oxford Pets):掉 1.5%-3%——需要纹理和细节
  • 纯纹理(DTD):掉 6.65%——最依赖高频信息,降采样伤害最大

结论:降采样有效与否,取决于任务所需的信息频段。不是所有分类任务都一样。

效率提升

输入尺寸 Patches Oxford Pets Acc 吞吐量
224×224 196 93.81% 791/s
168×168 100 90.65% 1458/s (+84%)
112×112 49 86.64% 2820/s (+257%)

隔行去行列(直接丢掉一半像素)比双线性降采样略差(低 0.5%),平滑插值保留更多信息。

灰度图实验(颜色信息剥离)

去掉颜色信息只保留亮度,ViT 还能分类吗?

数据集 Baseline (RGB) Grayscale 下降
Oxford Pets 93.81% 90.68% -3.13%

去掉颜色掉 3.13%,说明 Oxford Pets 上颜色信息有一定作用(部分品种靠毛色区分),但形状和纹理更重要。


但有个坑:不是所有任务都这么宽容

把 168×168 降采样放到多个数据集上:

数据集 特点 Baseline 168×168 下降幅度
CIFAR-100 粗粒度,原生 32×32 91.69% 91.56% -0.13%(假象)
Food-101 细粒度菜品 91.37% 89.87% -1.50%
Oxford Pets 猫狗品种,花纹纹理 93.81% 90.65% -3.16%
DTD 纯纹理分类 80.85% 74.20% -6.65%

(完整跨数据集对比见下方"总结"部分)


总结:跨数据集对比

不同数据集对降采样的敏感度差异巨大:

数据集 特点 Baseline 168×168 下降 112×112 下降
CIFAR-100 粗粒度,原生 32×32 91.69% 91.56% -0.13% ❌假象 90.00% -1.69%
Food-101 细粒度菜肴 91.37% 89.87% -1.50% 85.96% -5.41%
Oxford Pets 猫狗品种,花纹纹理 ~93.3% 90.65% -3.16% 86.64% -7.17%
DTD 纯纹理分类 80.85% 74.20% -6.65%

核心结论

  1. 降采样伤害取决于任务所需的信息类型

    • 粗粒度物体分类(CIFAR-100)→ 几乎不受影响(但也因为原生分辨率太低)
    • 细粒度分类(Food-101, Oxford Pets)→ 掉 1.5%~3%
    • 纹理分类(DTD)→ 掉 6.65%,最严重
    • 规律:越依赖高频细节的任务,降采样伤害越大
  2. 颜色信息有一定作用但不关键 —— Oxford Pets 灰度图掉 3.13%,部分品种靠毛色区分

  3. 学习型方案不如降采样简单有效

    • Oxford Pets 上对比(相同条件:batch=16, accum=2):
方法 Test Acc vs Baseline 额外成本
Baseline(全量微调) 91.91%
降采样 168x168 90.65% -1.26%
灰度图 90.68% -1.23%
MAE + 从头训练 Router 88.80% -3.11% 需 Router + MAE Decoder
MAE + 预热 Router 85.85% -6.06% 需蒸馏 + 训练,且有害

最简单的方案效果最好。


实验环境

  • 模型:ViT-B/16(85.8M 参数),ImageNet-21K 预训练
  • 数据集:CIFAR-100(100 类)、Food-101(101 类)
  • GPU:NVIDIA RTX 4090D(24GB)
  • 吞吐量测试:同一 GPU、batch=128、50 batches 取平均
  • 超参数:lr=3e-5, weight_decay=0.05, label_smoothing=0.1, CosineAnnealing 100 epochs

欢迎讨论和反馈。