Instructions to use 1999xia/ViT_Fast with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- timm
How to use 1999xia/ViT_Fast with timm:
import timm model = timm.create_model("hf_hub:1999xia/ViT_Fast", pretrained=True) - Notebooks
- Google Colab
- Kaggle
File size: 12,539 Bytes
54ee1eb | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 | <!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="utf-8">
<title>让 ViT 跑得更快</title>
<style>
body { max-width: 800px; margin: 40px auto; padding: 0 20px;
font-family: -apple-system, 'PingFang SC', 'Microsoft YaHei', sans-serif;
line-height: 1.8; color: #333; font-size: 16px; }
h1 { font-size: 28px; color: #111; border: none; margin-bottom: 10px; }
h2 { font-size: 22px; color: #222; border-bottom: 2px solid #eee;
padding-bottom: 8px; margin-top: 40px; }
table { border-collapse: collapse; margin: 15px 0; width: 100%;
font-size: 15px; }
th, td { border: 1px solid #ddd; padding: 10px 14px; text-align: center; }
th { background: #f5f5f5; font-weight: 600; }
code { background: #f0f0f0; padding: 2px 6px; border-radius: 3px;
font-size: 14px; }
blockquote { border-left: 4px solid #4a9eff; margin: 15px 0;
padding: 10px 20px; background: #f8faff; color: #555; }
hr { border: none; border-top: 2px solid #eee; margin: 30px 0; }
strong { color: #111; }
</style>
</head>
<body>
<h1>让 ViT 跑得更快:我们试了三种方案,最终发现最简单的最有效</h1>
<hr>
<h2>问题</h2>
<p>ViT-B/16 处理一张 224×224 的图片时,会切成 <strong>196 个 patch</strong>,全部送入 12 层 transformer。这是 <strong>33.85 GFLOPs</strong> 的计算量。</p>
<p>但直觉告诉我们:不是所有 patch 都同等重要。背景、天空、重复纹理——这些 patch 对分类贡献很小。如果能在进入 transformer 之前<strong>过滤掉不重要的 patch</strong>,就能省下大量计算。</p>
<p>于是我们开始了探索。方向很明确,但路走了三条。</p>
<hr>
<h2>网络架构</h2>
<p>在进入具体方案前,先理解模型的结构。</p>
<h3>基线:全量 ViT-B/16</h3>
<pre>
输入图片 (224×224)
↓
Patch Embedding (Conv2d, 16×16, stride=16) → 196 patches, 每个 768维
↓
Position Embedding + CLS Token
↓
12 层 Transformer Block(每层 = Attention + MLP + Residual)
↓
LayerNorm → 分类头 → 输出
</pre>
<h3>Patch Selection 架构(加了 Router)</h3>
<pre>
输入图片 (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 → 分类头 → 输出
</pre>
<p><strong>Router 的参数量(296K)仅为 ViT-B/16 整体(85.8M)的 0.34%</strong>,极小。</p>
<h3>Differentiable Top-K 原理</h3>
<p>选择"选 K 个 patch"是离散操作,没法直接算梯度。用 STE 技巧绕过:</p>
<pre>
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()
</pre>
<hr>
<h2>基线:全量 ViT-B/16 微调</h2>
<p>所有对比的基准是 <strong>ViT-B/16(85.8M 参数)在 ImageNet-21K 上预训练后,在目标数据集上全量微调</strong>。</p>
<p>训练配置:</p>
<ul>
<li>优化器:AdamW,lr=3e-5,weight_decay=0.05</li>
<li>调度器:CosineAnnealing,100 epochs(CIFAR-100)/ 30 epochs(Food-101)</li>
<li>数据增强:RandomResizedCrop + RandomHorizontalFlip</li>
<li>标签平滑:0.1</li>
<li>有效 batch size:128(batch=32 × accum=4)</li>
</ul>
<table>
<tr><th>数据集</th><th>类别</th><th>原生分辨率</th><th>Test Acc</th><th>适合降采样实验</th></tr>
<tr><td><strong>Oxford Pets</strong></td><td>37</td><td><strong>~200-1000px(高)</strong></td><td><strong>93.81%</strong></td><td>✅ 主数据集</td></tr>
<tr><td>Food-101</td><td>101</td><td>~300-512px(高)</td><td>91.37%</td><td>✅</td></tr>
<tr><td>DTD</td><td>47</td><td>~300-400px(高)</td><td>80.85%</td><td>✅ 纹理敏感</td></tr>
<tr><td>CIFAR-100</td><td>100</td><td>32x32(极低)</td><td>91.69%</td><td>❌ 不适合降采样</td></tr>
</table>
<p style="color:#999;font-size:14px"><strong>关于 CIFAR-100</strong>:原图仅 32x32,训练时放大到 224x224。本身就没多少细节,降采样测试会给出"几乎不掉精度"的假象。CIFAR-100 仅适合测试 patch 级别的选择策略(如 Router),不适合测试降采样/分辨率相关的实验。所有降采样实验应以 Oxford Pets、Food-101 等高分辨率数据集为准。</p>
<hr>
<h2>第一条路:Gumbel Selection —— 学不动</h2>
<p><strong>思路:</strong> 加一个小型 MLP Router(~300K 参数)给每个 patch 打分,用 Gumbel-Softmax + STE 做可微分的"选/不选",只把选中的 patch 送进 backbone。</p>
<p><strong>结果(CIFAR-100, 保留 50%):</strong></p>
<table>
<tr><th>方法</th><th>Acc</th><th>vs 全量</th></tr>
<tr><td>全量 ViT-B/16</td><td>91.69%</td><td>—</td></tr>
<tr><td>Gumbel Selection</td><td><strong>87.18%</strong></td><td><strong>-4.51%</strong></td></tr>
</table>
<p><strong>问题在哪?</strong> Gumbel 噪声太大了。Router 想表达"这个 patch 重要,score=0.8",但噪声一加,变成了随机数。梯度信号被淹没,Router 学了一个周期,最后还是没学会。</p>
<blockquote><strong>教训:</strong> 用噪声做离散化选择,噪声强度控制不好就白学。</blockquote>
<hr>
<h2>第二条路:MAE + 蒸馏 Router —— 有效,但复杂</h2>
<p><strong>思路:</strong> 去掉 Gumbel,改用 Sigmoid STE(梯度直通),让梯度直接穿过"选/不选"的门槛。同时加一个 <strong>MAE 重建任务</strong>:Encoder 处理后,用小 Decoder 重建被丢弃的 patch,用 MSE Loss 告诉 Router"你丢的东西能不能猜回来"。</p>
<p>这相当于给 Router 配了两个教练:</p>
<ul>
<li><strong>分类 Loss</strong> 告诉它:选出来的 patch 要能分类</li>
<li><strong>重建 Loss</strong> 告诉它:丢掉的 patch 不能是关键信息</li>
</ul>
<table>
<tr><th>方法</th><th>保留</th><th>Acc</th><th>vs 全量</th></tr>
<tr><td>全量 ViT-B/16</td><td>100%</td><td>91.69%</td><td>—</td></tr>
<tr><td>MAE + 随机 Router</td><td>50%</td><td>88.25%</td><td>-3.44%</td></tr>
<tr><td>MAE + 蒸馏 Router</td><td>50%</td><td><strong>89.07%</strong></td><td>-2.62%</td></tr>
<tr><td>MAE + 蒸馏 Router</td><td>75%</td><td><strong>91.10%</strong></td><td><strong>-0.59%</strong></td></tr>
</table>
<p>保留 75% 时吞吐量提升 27%,保留 50% 时提升 73%。</p>
<p><strong>但有个细节:</strong> 蒸馏 Router 的注意力重叠率只有 55.76%(随机基线 50%),几乎没学到。核心收益来自 MAE 而非蒸馏。</p>
<blockquote><strong>教训:</strong> 学习型方案能工作,但训练流程复杂(两阶段),且核心收益来自 MAE。</blockquote>
<hr>
<h2>第三条路:降采样训练 —— 简单到让人怀疑</h2>
<p><strong>"如果我只是把输入图片缩小一点,会怎么样?"</strong></p>
<pre>原图 224×224 → Resize(168×168) → Patch Embed → 100 patches → ViT → 分类</pre>
<p>不需要 Router、不需要蒸馏、不需要 MAE。就一行 Resize。</p>
<h3>但要小心:CIFAR-100 是个陷阱</h3>
<p>CIFAR-100 原图只有 32×32,放大到 224×224 训练本来就没细节。在 CIFAR-100 上测降采样得到"几乎不掉精度"的结论是假的——信息根本没有损失。</p>
<p>换到真正的高分辨率数据集上测试,结果完全不同:</p>
<table>
<tr><th>数据集</th><th>特点</th><th>Baseline</th><th>168×168</th><th>下降</th></tr>
<tr><td>CIFAR-100</td><td>粗粒度,原生 32×32</td><td>91.69%</td><td>91.56%</td><td>-0.13% (假象)</td></tr>
<tr><td><strong>Food-101</strong></td><td>细粒度菜品</td><td>91.37%</td><td><strong>89.87%</strong></td><td><strong>-1.50%</strong></td></tr>
<tr><td><strong>Oxford Pets</strong></td><td>猫狗品种,花纹重要</td><td>93.81%</td><td><strong>90.65%</strong></td><td><strong>-3.16%</strong></td></tr>
<tr><td><strong>DTD</strong></td><td>纯纹理分类</td><td>80.85%</td><td><strong>74.20%</strong></td><td><strong>-6.65%</strong></td></tr>
</table>
<p>精度下降幅度与数据集的信息依赖类型高度相关:粗粒度物体几乎不掉 → 细粒度掉 1.5-3% → 纯纹理掉最多(6.65%)。降采样有效与否,取决于任务所需的信息频段。</p>
<p><strong>效率提升(Oxford Pets):</strong></p>
<table>
<tr><th>输入尺寸</th><th>Patches</th><th>Acc</th><th>吞吐量</th></tr>
<tr><td>224×224</td><td>196</td><td>93.81%</td><td>791/s</td></tr>
<tr><td><strong>168×168</strong></td><td><strong>100</strong></td><td><strong>90.65%</strong></td><td><strong>1458/s (+84%)</strong></td></tr>
<tr><td>112×112</td><td>49</td><td>86.64%</td><td>2820/s (+257%)</td></tr>
</table>
<p>隔行去行列(直接丢掉一半像素)比双线性降采样略差(低 0.5%),平滑插值保留更多信息。</p>
<h3>灰度图实验(颜色信息剥离)</h3>
<p>去掉颜色信息只保留亮度,ViT 还能分类吗?</p>
<table>
<tr><th>数据集</th><th>Baseline (RGB)</th><th>Grayscale</th><th>下降</th></tr>
<tr><td>Oxford Pets</td><td>93.81%</td><td><strong>90.68%</strong></td><td><strong>-3.13%</strong></td></tr>
</table>
<p>去掉颜色掉 3.13%,说明 Oxford Pets 上颜色信息有一定作用(部分品种靠毛色区分),但形状和纹理更重要。</p>
<hr>
<h2>但有个坑:不是所有任务都这么宽容</h2>
<table>
<tr><th>数据集</th><th>全量</th><th>168×168</th><th>下降幅度</th></tr>
<tr><td>CIFAR-100</td><td>91.69%</td><td>91.56%</td><td><strong>-0.13%</strong></td></tr>
<tr><td>Food-101</td><td>91.37%</td><td>89.87%</td><td><strong>-1.50%</strong></td></tr>
</table>
<p>Food-101 掉点多了 10 倍。CIFAR-100 的物体靠轮廓区分,Food-101 的食物靠纹理区分。细粒度分类对分辨率更敏感。</p>
<blockquote>分类粒度越细,对分辨率越敏感。粗粒度适合降采样,细粒度要谨慎。</blockquote>
<hr>
<h2>如果推理时缩小(不重新训练)呢?</h2>
<p>效果很差:</p>
<ul>
<li>168×168 推理硬切:89.24%(vs 重新训练 91.56%)</li>
<li>112×112 推理硬切:73.71%(vs 重新训练 90.00%)</li>
</ul>
<p>ViT 的位置编码在训练时学会的,推理时改变输入尺寸对齐全乱。<strong>必须重新训练。</strong></p>
<hr>
<h2>总结:跨数据集对比</h2>
<table>
<tr><th>数据集</th><th>特点</th><th>Baseline</th><th>168×168</th><th>下降</th><th>112×112</th><th>下降</th></tr>
<tr><td>CIFAR-100</td><td>粗粒度,32×32原生</td><td>91.69%</td><td>91.56%</td><td>-0.13%</td><td>90.00%</td><td>-1.69%</td></tr>
<tr><td>Food-101</td><td>细粒度菜肴</td><td>91.37%</td><td>89.87%</td><td>-1.50%</td><td>85.96%</td><td>-5.41%</td></tr>
<tr><td><strong>Oxford Pets</strong></td><td>猫狗品种,花纹纹理</td><td><strong>~93.3%</strong></td><td><strong>90.65%</strong></td><td><strong>-3.16%</strong></td><td><strong>86.64%</strong></td><td><strong>-7.17%</strong></td></tr>
<tr><td>DTD</td><td>纯纹理分类</td><td>80.85%</td><td>74.20%</td><td>-6.65%</td><td>—</td><td>—</td></tr>
</table>
<h3>核心结论</h3>
<ol>
<li><strong>降采样伤害取决于任务所需的信息类型</strong>——粗粒度几乎不受影响,细粒度掉 1.5-3%,纹理分类掉最多(6.65%)。越依赖高频细节的任务,降采样伤害越大。</li>
<li><strong>颜色信息有一定作用但不关键</strong>——Oxford Pets 灰度图掉 3.13%,部分品种靠毛色区分。</li>
<li><strong>学习型方案 vs 降采样</strong>——待跑完随机/蒸馏 Router 对比后再下结论。</li>
</ol>
<hr>
<p style="color: #999; font-size: 14px;">欢迎讨论和反馈。</p>
</body>
</html>
|