timm
File size: 11,534 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
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
# 让 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

---

*欢迎讨论和反馈。*