timm
File size: 15,341 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
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
# Miss Patch — Pre-Tokenization Patch Selection for Vision Transformers

在 Vision Transformer(ViT)前插入一个 Router/Filter 模块,动态过滤不重要的 image patches,减少计算量,同时尽量保持分类精度。

---

## 项目文件

| 文件 | 说明 |
|------|------|
| `models.py` | 所有模型定义,见下方"模型类"表格 |
| `train.py` | 通用训练脚本(支持 baseline、patch selection、MAE finetune 等) |
| `train_patch_selection_mae.py` | MAE patch selection 专用训练脚本(含蒸馏 router 加载) |
| `train_router_distill.py` | Attention Distillation 训练脚本 |
| `test_patch_selection_b16.py` | 旧 Gumbel patch selection 测试脚本 |
| `test_blur_downsample.py` | 图片模糊/降采样测试脚本 |
| `test_stride_patches.py` | 不同 stride 下 patch 数量 vs 精度测试脚本 |
| `datasets.py` | 数据集加载器:CIFAR-10/100, Oxford Pets, Food-101, Tiny-ImageNet, DTD, Flowers-102, Stanford Cars |
| `checkpoints/` | 所有模型权重 |
| `logs/` | 训练日志 |
| `SHARED_MEMORY.md` | 跨 Claude Code 会话共享记忆 |

### models.py — 模型类

| 类名 | 说明 |
|------|------|
| `GumbelSelection` | Gumbel-Softmax 可微分选择模块(含 STE + 退火温度) |
| `SemanticRouter` | 语义 Router(MLP: D→D/2→1 + GELU + LayerNorm) |
| `PatchSelectionViT` | Patch Selection ViT(Router + 可微分 Top-K + Gumbel 噪声) |
| `RandomPruneViT` | 随机丢弃 patch 的 baseline |
| `MAEDecoder` | MAE 解码器(4 blocks, 512-dim transformer) |
| `MAEViT` | 标准 MAE(mask 75% patches + 重建) |
| `MAEPatchSelectionViT` | **MAE + Patch Selection**(Router + Differentiable Top-K + 轻量 encoder + 10 blocks backbone + MAE Decoder) |

### models.py — create_model() 可用 model_name

| model_name | 对应模型 | 说明 |
|-----------|---------|------|
| `swin_tiny` | Swin-Tiny | Swin Transformer baseline |
| `patch_selection_vit` | PatchSelectionViT | Gumbel 方案(旧) |
| `patch_selection_vit_b16` | PatchSelectionViT | ViT-B/16 IN-21K + Gumbel |
| `patch_selection_vit_b16_in1k` | PatchSelectionViT | ViT-B/16 IN-1K + Gumbel |
| `random_prune_vit` | RandomPruneViT | 随机丢弃 50% baseline |
| `mae_vit` | MAEViT | 标准 MAE 预训练 |
| `mae_patch_selection_vit_b16` | MAEPatchSelectionViT | **MAE + 蒸馏 Router(推荐)** |

---

## 使用方法

### 1. 训练 Baseline(全量 ViT-B/16)

```bash
# CIFAR-100 (IN-21K pretrained)
python train.py --model vit_b16 --dataset cifar100 --lr 3e-5 --epochs 100 --batch_size 128

# Oxford Pets
python train.py --model vit_b16 --dataset oxford_pets --lr 3e-5 --epochs 100 --batch_size 32

# Food-101
python train.py --model vit_b16 --dataset food101 --lr 3e-5 --epochs 30 --batch_size 32
```

### 2. 训练 MAE Patch Selection(+ 蒸馏 Router)

**两步走:**

**Step 1:** 训练蒸馏 Router(从教师 ViT 学习注意力分布)
```bash
python train_router_distill.py --dataset cifar100 --gpu 4
# 可选参数:
#   --batch_size 64     # 批量大小(默认 64)
#   --lr 1e-4           # 学习率(默认 1e-4)
#   --epochs 5          # 训练轮数(默认 5)
#   --weight_decay 0.05 # 权重衰减

# 输出:checkpoints/router_distill_{dataset}/router.pth
```

**Step 2:** 用蒸馏 Router 初始化,训练完整模型
```bash
python train_patch_selection_mae.py --dataset cifar100 --gpu 4 \
  --router_path ./checkpoints/router_distill_cifar100/router.pth

# 可选参数:
#   --keep_ratio 0.5    # 保留 patch 比例(默认 0.5,推荐 0.75)
#   --batch_size 32     # 批量大小(默认 32)
#   --accum 4           # 梯度累积步数(默认 4,有效 batch = 128)
#   --lr 3e-5           # 学习率(默认 3e-5)
#   --weight_decay 0.05 # 权重衰减
#   --label_smoothing 0.1
#   --mse_start 1.0     # MSE 损失起始权重
#   --mse_end 0.1       # MSE 损失结束权重(cosine anneal)
#   --decoder_dim 512   # MAE decoder 维度
#   --decoder_depth 4   # MAE decoder 层数
#   --router_path PATH  # 蒸馏 router 权重路径(可选)
```

**不加 `--router_path` 就是随机初始化 Router**(原始 MAE 方案)。

### 3. 训练 Gumbel Patch Selection(旧方案,不推荐)

```bash
python train.py --model patch_selection_vit_b16 --dataset cifar100
```

### 4. 测试图片预处理(模糊/降采样)

```bash
python test_blur_downsample.py --dataset cifar100 --gpu 5
# 测试 Gaussian blur、降采样等预处理对精度的影响(不减少 token 数量)
```

### 5. 测试 Stride-based Patch 减少

```bash
python test_stride_patches.py --dataset cifar100 --gpu 5
# 测试不同 stride 下 patch 数量 vs 精度(直接减少 token 数量)
```

### 6. 降采样图片训练(减小输入分辨率)

```bash
# 112×112 → 49 patches (75% reduction)
python test_downsample_train.py --dataset cifar100 --image_size 112 --gpu 3

# 168×168 → 100 patches (49% reduction)
python test_downsample_train.py --dataset cifar100 --image_size 168 --gpu 5

# 可选参数:
#   --batch_size 128   # 批量大小(默认 128)
#   --lr 3e-5          # 学习率(默认 3e-5)
#   --epochs 100       # 训练轮数(默认使用数据集默认值)
#   --label_smoothing 0.1
#   --weight_decay 0.05
#   --num_workers 4

# 输出:checkpoints/{dataset}_vit_b16_img{size}/best_model.pth
```

### 7. 通用训练脚本 train.py 全部参数

```bash
python train.py \
  --model vit_b16                # 模型名
  --dataset cifar100             # 数据集
  --batch_size 128               # 批量大小
  --epochs 100                   # 训练轮数
  --lr 3e-5                      # 学习率
  --weight_decay 0.05            # 权重衰减
  --keep_ratio 0.5               # patch 保留比例
  --selection_mode topk          # 选择模式
  --adaptive_alpha 0.5           # 自适应 alpha
  --patch_size 16                # patch 大小
  --patch_stride 16              # patch stride
  --num_workers 10               # 数据加载线程
  --use_randaugment              # 使用 RandAugment
  --label_smoothing 0.0          # 标签平滑
  --device auto                  # 设备
  --save_dir ./checkpoints       # 保存目录
  --pretrained                   # 使用预训练权重
  --load_mae PATH                # 加载 MAE 预训练权重
  --log_dir ./logs               # 日志目录
  --drop_path 0.0                # DropPath 率
  --mixup 0.0                    # MixUp alpha
  --cutmix 0.0                   # CutMix alpha
  --ema_decay 0.0                # EMA 衰减率
  --clip_grad 1.0                # 梯度裁剪
  --warmup_epochs 0              # 预热轮数
```

### 支持的数据集

`datasets.py` 提供:
- `get_cifar10_loader` — CIFAR-10 (10 classes, 50K train)
- `get_cifar100_loader` — CIFAR-100 (100 classes, 50K train)
- `get_oxford_pets_loader` — Oxford Pets (37 classes, ~3.6K train)
- `get_food101_loader` — Food-101 (101 classes, 68K train)
- `get_tiny_imagenet_loader` — Tiny-ImageNet (200 classes, 100K train)
- `get_dtd_loader` — DTD (47 classes, texture)
- `get_flowers102_loader` — Flowers-102 (102 classes)
- `get_stanford_cars_loader` — Stanford Cars (196 classes)

---

## 项目架构

### Baseline 架构(全量 ViT-B/16 IN-21K)

```
Image (224×224) → Patch Embed (196 patches, 768-dim) → 12 ViT-B blocks → CLS → Classifier
```

### MAE Patch Selection 架构(当前最佳)

```
Image → Patch Embed + Pos Embed → Router (MLP 768→384→1) → Sigmoid STE Top-K (keep 75%)
  → 选中 patches (147个) + CLS → 轻量 Encoder (前 2 个 ViT-B block)
    → 主干 (后 10 个 ViT-B block) → CLS → CE Loss
    → MAE Decoder (4 blocks, 512-dim) → 重建丢弃 patches → MSE Loss (λ: 1.0→0.1)
```

### Attention Distillation 流程

```
教师 ViT-B/16(冻结)→ 提取最后 block 的 CLS→patch 注意力 → 注意力分数 (B, 196)
MLP Router(学生)→ 从 patch embedding 预测重要性分数 → MSE Loss
→ 训练好的 Router 权重→ 初始化 MAEPatchSelectionViT → 端到端微调
```

---

## 实验结果汇总

### 主数据集:Oxford-IIIT Pets(37 类,高分辨率 ~200-1000px)

**全量对比(相同条件:batch=16, accum=2, lr=3e-5, 100 epochs):**

| 方法 | 保留 | Test Acc | vs Baseline |
|:----|:----:|:--------:|:----------:|
| **Baseline(全量 ViT-B/16)** | 100% | **91.91%** | — |
| 降采样 168x168 | 51% patches | **90.65%** | **-1.26%** |
| 灰度图(Grayscale) | 亮度 only | **90.68%** | **-1.23%** |
| **MAE + 从头训练 Router** | 50% patches | **88.80%** | **-3.11%** |
| 降采样 112x112 | 25% patches | **86.64%** | **-5.27%** |
| **MAE + 预热 Router** | 50% patches | **85.85%** | **-6.06%** |

**结论:**
- 降采样 168x168 和灰度图效果最好(掉 ~1.2%),且不需要额外模块
- MAE + 从头训练 Router(88.80%)比预热过的 Router(85.85%)更好——在小数据集上,教师蒸馏信号太弱(重叠率仅 51.16%≈随机),预热反而有害
- 降采样比所有学习型方案更简单、效果更好

---

### 验证数据集

#### Food-101(101 类,细粒度菜肴,~300-512px)

| 方法 | Test Acc | vs Baseline |
|:----|:--------:|:----------:|
| Baseline(batch=128) | 91.37% | — |
| 降采样 168x168 | 89.87% | -1.50% |
| 降采样 112x112 | **85.96%** | -5.41% |
| MAE + 预热 Router(50%,CIFAR) | 89.52% | -1.85%(参考) |

#### DTD(47 类,纹理分类,~300-400px)

| 方法 | Test Acc | vs Baseline |
|:----|:--------:|:----------:|
| Baseline | 80.85% | — |
| 降采样 168x168 | **74.20%** | **-6.65%** |

---

### 跨数据集对比:降采样 168x168

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

**规律:分类粒度越细、越依赖高频纹理,降采样伤害越大。**

---

### 旧实验(CIFAR-100,仅供参考)

> CIFAR-100 原生仅 32x32,降采样结果不可靠,但 Router 类实验(不改变图片大小)有效。

| 方法 | 保留 | CIFAR-100 | vs BL |
|:----|:----:|:--------:|:-----:|
| Baseline | 100% | 91.69% | — |
| Gumbel Selection | 50% | 87.18% | -4.51% |
| MAE + 随机 Router | 50% | 88.25% | -3.44% |
| MAE + 蒸馏 Router | 50% | 89.07% | -2.62% |
| MAE + 蒸馏 Router | 75% | 91.10% | -0.59% |## 关键发现

1. **Gumbel 方案最差** — Gumbel 噪声导致 Router 梯度几乎消失,无法有效学习 patch 重要性
2. **MAE 重建损失有帮助** — 让 Router 通过"能否重建丢弃 patch"来学习,比纯分类信号好(+1%)
3. **Attention Distillation 最佳** — 用教师 ViT 的 CLS 注意力预训练 Router,再端到端微调
   - CIFAR-100: 88.25% → **89.07%**(+0.82%)
   - Food-101: ~87.97% → **89.52%**(+1.55%)
4. **数据集大小重要** — Oxford Pets 仅 ~3.6K 训练图像,蒸馏信号太弱,反而下降
5. **CLS token flow 关键 bug** — 早期实现中 CLS 只走 10/12 block(前 2 block 浪费),修复后 epoch 1 从 80.63%→84.69%
6. **75% keep ratio 是更优平衡点** — 只掉 0.59%(vs 50% 掉 2.62%),推荐使用
7. **降低 token 质量 vs 减少 token 数量** — 模糊/降采样图片几乎不影响精度(-0.14~-0.57%),但计算量完全不变;减少 token 数量才真正降低计算量
8. **Naive stride 减少 token 效果差且实用性低** — stride=32 减少 75% token 时掉 21%,远不如学习型 Router 或降采样训练

---

## 检查点文件

```
checkpoints/
├── cifar100_mae_patchsel_b16_keep50/
│   └── best_model.pth          —— 89.07% (MAE + 蒸馏 Router, keep 50%)
├── cifar100_mae_patchsel_b16_keep75/
│   └── best_model.pth          —— 91.10% (MAE + 蒸馏 Router, keep 75%) ★推荐
├── food101_mae_patchsel_b16_keep50/
│   └── best_model.pth          —— 89.52% (MAE + 蒸馏 Router, keep 50%)
├── oxford_pets_mae_patchsel_b16_keep50/
│   └── best_model.pth          —— 85.99% (MAE + 蒸馏 Router, keep 50%)
├── cifar100_patchsel_b16_keep50/
│   └── best_model.pth          —— 87.18% (Gumbel, CIFAR-100)
├── food101_patchsel_b16_keep50/
│   └── best_model.pth          —— 85.77% (Gumbel, Food-101)
├── oxford_pets_patchsel_b16_keep50/
│   └── best_model.pth          —— 87.00% (Gumbel, Oxford Pets)
├── cifar100_random_prune_vit_keep50_topk/
│   └── best_model.pth          —— 48.13% (随机丢弃 50%)
├── cifar100_vit_b16_ft/
│   └── best_model.pth          —— 88.75% (全量微调, CIFAR-100)
├── cifar100_vit_small/
│   ├── best_model.pth          —— 84.83% (ViT-S/16 baseline)
│   └── checkpoint_epoch*.pth   —— 每 10 epoch 检查点
├── cifar100_patch_selection_vit_keep50_topk/
│   ├── best_model.pth          —— 73.01% (旧 Gumbel 实验)
│   └── checkpoint_epoch*.pth   —— 每 10 epoch 检查点
├── vit_b16_cifar10_best.pth    —— CIFAR-10 baseline
├── vit_b16_cifar10_ft/
│   └── best_model.pth          —— CIFAR-10 微调
├── dtd_vit_b16_in21k/
│   └── best_model.pth          —— DTD 80.85%
├── flowers102_vit_b16_in21k/
│   └── best_model.pth          —— Flowers-102 100.00%
├── food101_vit_b16_in21k/
│   └── best_model.pth          —— Food-101 86.68% (val)
├── oxford_pets_vit_b16_ft/
│   └── best_model.pth          —— Oxford Pets 全量微调
├── oxford_pets_vit_b16_in21k/
│   └── best_model.pth          —— Oxford Pets IN-21K 95.65% (val)
├── mae_cifar100_mask75/
│   ├── mae_best.pth            —— MAE 预训练最佳
│   ├── mae_encoder_final.pth   —— MAE encoder 最终
│   ├── mae_epoch50.pth         —— 50 epoch 检查点
│   └── mae_epoch100.pth        —— 100 epoch 检查点
├── router_distill_cifar100/
│   └── router.pth              —— 蒸馏 Router (CIFAR-100, 55.76% overlap)
├── router_distill_food101/
│   └── router.pth              —— 蒸馏 Router (Food-101, 59.95% overlap)
└── router_distill_oxford_pets/
    └── router.pth              —— 蒸馏 Router (Oxford Pets, 51.16% overlap)
```

---

## 超参数(推荐)

所有 MAE patch selection 实验统一使用:

| 参数 | 值 |
|:----|:---|
| Backbone | ViT-B/16 IN-21K pretrained (`vit_base_patch16_224.augreg_in21k`) |
| Batch size | 32 |
| Gradient accumulation | 4(effective batch = 128) |
| Learning rate | 3e-5 |
| Weight decay | 0.05 |
| Label smoothing | 0.1 |
| Scheduler | CosineAnnealingLR |
| Gradient clip | 1.0 |
| Keep ratio | **0.75**(147/196 patches,推荐)或 0.5(98/196 patches) |
| MSE weight | 1.0 → 0.1(cosine anneal) |
| Decoder dim | 512 |
| Decoder depth | 4 |
| Epochs | CIFAR-100: 100, Oxford Pets: 100, Food-101: 30 |

---

*最后更新: 2026-05-04*