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