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: 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*
|