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