# 让 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 技巧绕过: ```python 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 --- *欢迎讨论和反馈。*