简介:FastVIT实战:使用FastVIT实现图像分类是一份面向图像分类初学者与Transformer研究者的完整实践资源,以轻量高效的FastVIT架构为核心,系统演示从数据预处理、数据增强、模型训练、验证到导出与测试的全流程。压缩包共含2000个文件,以1979张png样本图片为主,配合10个py脚本、2个json配置文件、pt/pth模型权重以及少量pyc缓存和txt说明,整体大小约764.79MB,目录结构清晰,便于按步骤复现。资源不仅覆盖图像分类的基础概念,还分别拆解makedata.py、train.py、export_model.py、test.py等脚本的核心职责:如何构建数据集与数据加载器,如何初始化FastVIT模型并设置损失函数与优化器,如何完成训练与验证循环,以及如何导出权重并在测试集上评估泛化能力。已有648人学习,适合希望快速跑通Transformer图像分类项目、并深入理解模型训练与部署细节的学习者。
1. FastVIT 实战要解决的是哪种图像分类问题
同样做图像分类,最新的图像分类模型在榜单上很好看,一放进产品里就现原形:显存不够、延迟超预算、模型文件太大。FastVIT 的价值不是让我在 ImageNet 上再刷一个点,而是让推理时间真正压下来。它用卷积和自注意力混合的结构,把注意力算子放到最划算的位置,训练时多分支、推理时折算成单分支卷积,所以精度和速度能同时站住。
这篇笔记的目标读者很直接:手头有自定义图像分类数据集,想从 CNN 换到 ViT,又不想被部署成本卡死的开发者和算法工程师。下面按落地路径写:数据怎么准备、训练参数怎么设、坑在哪、导出要注意什么,都是可复现的配置,不是拿模型跑个 demo 就结束。
2. FastVIT 为什么适合做图像分类:卷积与注意力的混合结构
2.1 从 ViT 的部署痛点到 FastVIT 的混合设计
很多图像分类算法教程上来就让你看 ViT 结构,但真正落地时最大的问题不是 ViT 不 work,而是延迟不达标。ViT 把 224x224 的图切成 patch,变成 196 个 token,前几层就开始做全局自注意力。每个 token 要和另外 195 个 token 算相似度,十二层叠下来,计算量和访存量马上失控。
问题在于,图像分类的前期特征主要是边缘、纹理、颜色渐变这些局部信息,并不需要在一开始就看到整张图。CNN 用一个 3x3 卷积就能把局部关系学好,而且参数少、硬件友好。FastVIT 的做法是别把 Transformer 和 CNN 对立起来,而是把特征提取的前半段交给卷积,后半段再交给注意力。
FastVIT 前期的 block 里用了一个叫 RepMixer 的混合算子,数学上可以理解成一个局部 mixing 操作,训练完以后能融合成一个标准卷积。后面的 stage 才用自注意力做整图信息的聚合。这样设计以后,高分辨率、大 token 数的前期阶段避免了一次昂贵注意力计算,后期特征图分辨率已经变小,注意力开销也就可控了。
我自己拿纯 CNN 和纯 ViT 跑过同一份自定义数据集。纯 CNN 在小数据集上训练稳定,但到了复杂背景的森林图像分类这种场景,远处目标和背景之间的关系抓不住;纯 ViT 在 ImageNet 预训练加持下能学到,但推理时贵。FastVIT 正好夹在中间:前段用卷积学纹理,后段用注意力学关系,这也是它在图像分类任务上做工整的原因。
2.2 训练态与推理态分离:结构重参数化到底省在哪
FastVIT 最容易被忽略的一点是:它不是一个训练完直接能跑的模型,而是有两种状态。训练态里,RepMixer 和部分卷积块保留了多分支结构。常见做法是残差分支、卷积分支、BN 分支同时存在,梯度可以从多条路径回传,优化更稳。
真正有意思的是推理态。训练结束后,多分支可以按数学等价关系折算成单分支卷积。BN 的 scale 和 shift 融进卷积权重,残差分支折算成 1x1 卷积或者直接加到卷积核中心,最终只留下一个卷积分支。
这个操作不改变模型输出,但省掉的是推理时大量分支之间的 upsample、add、BN 算子和额外的内存搬运。算子少了,CPU 和 GPU 上的延迟都会明显下降。移动端和边缘设备上尤其明显,因为这类硬件对算子数量很敏感,一个小分支就可能多触发一次 kernel launch。
这里有个容易翻车的习惯:以为model.eval()就是推理态。eval 只改变 BN 和 Dropout 的行为,不合并权重。如果你直接拿训练态的权重去转 ONNX,模型里会残留大量分支结构,导出文件变大,推理也慢。后面第 5 章我会专门说这个坑。
2.3 FastVIT 与 ResNet/EfficientNet 的选型边界
FastVIT 不是万能模型,我一般会给团队一个很粗的选型标准。如果项目部署在数据中心 GPU 上,延迟要求不苛刻,ResNet 或者 EfficientNet 更省心,生态和预训练权重都成熟。如果要做端侧实时图像分类,而且对准确率还有要求,FastVIT 这类混合模型就值得试。
具体到图像分类任务,还要看类别之间的差异。细粒度分类,比如森林图像分类里要区分相似树种,注意力机制对全局形状关系有帮助,FastVIT 会比同量级 MobileNet 更稳。如果只是区分猫狗这种大类,MobileNetV3 就够了,没必要引入结构重参数化这条额外链路。
FastVIT 的代价是工程上比 ResNet 多一步:你要会管理训练态和推理态,要懂得在导模型前做分支融合。这个复杂度换来的是推理延迟的明显下降,到底值不值,取决于你的部署环境是不是真的卡在延迟上。我的建议是先用 ImageNet 预训练权重在你的测试集上跑一遍,再决定要不要全面替换现有图像分类模型,别一上来就重构。
3. 图像分类数据集准备:目录规范、标签划分与增强配置
3.1 目录规范:用 ImageFolder 还是自建 Dataset
训练脚本里最常见的数据加载方式是用torchvision.datasets.ImageFolder,前提是目录结构按照类别分好。FastVIT 对输入本身不挑数据集格式,PyTorch 能读什么它就能训什么,所以目录规范越早定好,后面越省事。
我一般会把项目数据整理成下面这个结构:
data/my_dataset/ train/ forest/ forest_001.jpg forest_002.jpg river/ river_001.jpg village/ village_001.jpg val/ forest/ forest_010.jpg river/ river_003.jpg创建目录用一条 bash 命令就能完成。类别名不要用中文,也不要用带空格的目录名,否则跨服务器拷贝和后续脚本处理都可能出问题。
DATA_ROOT="data/my_dataset" mkdir -p "$DATA_ROOT"/train "$DATA_ROOT"/val # 按你的类别列表生成目录 for c in forest river village city; do mkdir -p "$DATA_ROOT/train/$c" "$DATA_ROOT/val/$c" done这样做的理由是ImageFolder会按目录名排序后生成 label,类别顺序不是你塞数据的顺序,而是字符串排序后的顺序。如果第 3 个类别实际是village,但代码里写死 label 2 是forest,验证集准确率照样能看,线上全错。所以后续要保存一份class_to_idx的 JSON,不要靠记忆。
如果原始图片是乱七八糟放在一个大目录里的,我建议先写一个划分脚本,而不是手工拖文件。下面这个脚本会按比例随机划分,并尽量保持每个类别的样本比例一致。
import os import random import shutil from glob import glob random.seed(42) train_ratio = 0.85 src_root = "raw/all" dst_root = "data/my_dataset" # 自动发现一级子目录作为类别 classes = sorted([d for d in os.listdir(src_root) if os.path.isdir(os.path.join(src_root, d))]) for cls in classes: cls_dir = os.path.join(src_root, cls) images = [] for ext in ("*.jpg", "*.jpeg", "*.png"): images.extend(glob(os.path.join(cls_dir, ext))) random.shuffle(images) split_point = int(len(images) * train_ratio) train_images = images[:split_point] val_images = images[split_point:] for img in train_images: target = os.path.join(dst_root, "train", cls, os.path.basename(img)) os.makedirs(os.path.dirname(target), exist_ok=True) shutil.copy2(img, target) for img in val_images: target = os.path.join(dst_root, "val", cls, os.path.basename(img)) os.makedirs(os.path.dirname(target), exist_ok=True) shutil.copy2(img, target)脚本里有三个参数要按项目改:train_ratio、random.seed、扩展名列表。train_ratio在数据量少于一万张时我通常给 0.85,数据量大可以放宽到 0.9。random.seed固定下来,方便别人复现那次实验。
复制用shutil.copy2而不是shutil.move,是因为原始文件一旦移动坏了很麻烦,我吃过这个亏。如果你对数据备份有信心,也可以改成os.link做硬链接,省空间且速度快,前提是原始目录和项目目录在同一个文件系统上。
3.2 数据划分:train/val 分开放,而不是所有文件随机丢
很多人会把所有图片放在同一个目录,然后用一个 CSV 记录哪张属于训练集哪张属于验证集。这种做法不是不行,但对 FastVIT 这种需要大量增广和随机采样的训练流程,ImageFolder的方式更省心,也少一层索引逻辑。
分训练集和验证集时,有一类风险是“同源数据泄漏”。比如森林图像分类里,同一棵树在不同角度拍了好几张,这些图如果一部分进了训练集、一部分进了验证集,验证准确率会虚高。严格的做法是先把图片按拍摄场景、地块或者视频片段分组,再对组做划分,而不是对单张图片做划分。代码逻辑上,只把group这个字段加入随机洗牌的单位即可。
如果数据真的少,比如每个类别只有三五百张,我倾向先不做复杂的训练验证划分,改成五折交叉验证,用折间均值和方差来判断一个图像分类模型是否稳定。单次划分很容易因为某几张小图让结果忽高忽低,这会把调参过程带偏。
验证集里每个类别的图片数量最好差不多。你想判断模型整体精度,类别不均衡会导致 val 分数全被大类带跑。我一般会在划分脚本最后统计每个类别的数量,打印出来看一眼,再开始训练。
3.3 增强和归一化:FastVIT 对预训练输入分布很敏感
FastVIT 用的是 ImageNet 预训练权重,输入分布默认是 ImageNet 的 mean 和 std。直接用ToTensor()然后喂给模型,精度会低一截,这不是模型问题,是输入分布没对上。
我在代码里一般直接调用timm.data.create_transform,它会把归一化和增强一起处理掉。
from timm.data import create_transform # 训练集增强 train_transform = create_transform( input_size=224, is_training=True, color_jitter=0.4, auto_augment="randaug", re_prob=0.25, ) # 验证集只做 resize、crop 和归一化 val_transform = create_transform( input_size=224, is_training=False, mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.255), )create_transform在is_training=True时默认会用 ImageNet 的 mean/std,还会按模型配置补上合适的 resize 策略。re_prob=0.25是 RandomErasing 的概率,相当于随机遮挡一部分区域,对森林图像分类这种背景复杂的场景很有效,可以减少模型只靠纹理斑块判断类别。
如果你是手写 transform,建议至少包含RandomResizedCrop(224)、RandomHorizontalFlip()、ColorJitter(0.4)三件套。CutMix 和 Mixup 这种重增广放到训练脚本里用timm.data.Mixup做,不要在 transform 里手工实现,否则 batch 维度处理容易出 bug。
有一个容易忽略的细节:验证集不要用RandomResizedCrop,它会让同一张图每次验证结果都不同。验证时固定CenterCrop(224)或者Resize(224),保证指标可复现。我一般会在验证集上只跑一次,而不是反复验证,因为模型训练过程中的随机性已经够多了。
4. 用 FastVIT 训练图像分类模型:配置、命令与调参记录
4.1 用 timm 加载预训练 FastVIT 并替换分类头
FastVIT 官方仓库的训练代码基于 PyTorch 和 timm,所以我在自己的项目里也用 timm 来加载模型,省去自己拼网络结构的时间。
import timm import torch NUM_CLASSES = 10 device = "cuda" if torch.cuda.is_available() else "cpu" model = timm.create_model( "fastvit_t8", pretrained=True, num_classes=NUM_CLASSES, ) model.to(device)pretrained=True会加载 ImageNet-1k 预训练权重。FastVIT 的模型命名一般带t8、t12这类后缀,代表不同深度,t8适合快速验证,t12适合精度优先的场景。第一次运行时权重会下载到本地缓存,内网机器要提前把权重准备好,否则程序卡在下载这一步。
timm.create_model传了num_classes之后,如果和预训练模型的 1000 类不一致,分类头会被重置。这是理所当然的,但很多人看到最后全连接层参数被随机初始化就慌,其实这正是迁移学习的标准流程。
如果你是自己写训练脚本,加载完模型后最好打印一行分类头的维度确认一下:
print(model.get_classifier())不要等到训练到一半报 shape mismatch 再回头查。这个打印动作花不了两秒钟,却能省掉很多无意义的排错时间。
4.2 训练主循环里的关键参数:优化器、EMA、AMP
FastVIT 我在微调时用的优化器是 AdamW,学习率 1e-3 左右,权重衰减 0.05,配合 cosine 学习率下降。训练轮数取决于数据量,公开数据集的完整训练要 300 轮以上,但我们做自定义图像分类微调,通常 30 到 50 轮就能看到收敛趋势。
下面是一个我常用的训练主循环骨架:
import torch.nn as nn from timm.optim import AdamW from timm.scheduler import CosineLRScheduler from timm.utils import ModelEmaV2 EPOCHS = 50 BATCH_SIZE = 64 criterion = nn.CrossEntropyLoss(label_smoothing=0.1) optimizer = AdamW(model.parameters(), lr=1e-3, weight_decay=0.05) ema_model = ModelEmaV2(model, decay=0.9998) scheduler = CosineLRScheduler( optimizer, t_initial=EPOCHS, warmup_t=5, warmup_lr_init=1e-5, ) scaler = torch.cuda.amp.GradScaler() for epoch in range(EPOCHS): model.train() for x, y in train_loader: x, y = x.to(device), y.to(device) optimizer.zero_grad() with torch.cuda.amp.autocast(): logits = model(x) loss = criterion(logits, y) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0) scaler.step(optimizer) scaler.update() ema_model.update(model) scheduler.step(epoch)重点说几个参数。
label_smoothing=0.1可以防止模型对训练标签过于自信,尤其适合类别之间有重叠的数据集。森林图像分类里,有些类别的纹理本来就像,硬标签会把边界学得过于尖锐。
ema_model.decay=0.9998是 EMA 的衰减系数。EMA 相当于对权重做了滑动平均,能减少后期训练震荡。验证的时候用ema_model.module而不是原始模型,经验上准确率通常更高更稳。
clip_grad_norm_的max_norm=5.0是梯度裁剪,配合 AMP 使用能规避一部分梯度爆炸。如果你用的是 PyTorch 2.x,torch.cuda.amp.autocast可以替换成torch.amp.autocast("cuda"),逻辑完全一样。
下面这个参数表是我微调 FastVIT 时的默认起点,实际项目里按数据量微调:
| 参数 | 推荐值 | 说明 |
|---|---|---|
| 优化器 | AdamW | ViT 系列默认选择 |
| 初始学习率 | 1e-3 | 小数据集降到 5e-4 |
| 权重衰减 | 0.05 | 正则化主力 |
| batch size | 32 ~ 128 | 按显存调整 |
| warmup epochs | 5 | 避免前期震荡 |
| EMA decay | 0.9998 | 验证用 EMA 权重 |
| label smoothing | 0.1 | 类别相似时效果好 |
4.3 日志与验证指标:怎么判断真的训好了
训练过程中要同时盯三个指标:训练 loss、验证 top-1 准确率、验证 loss。只盯准确率很容易被一两个 epoch 的波动骗到。FastVIT 前几轮由于分类头刚随机初始化,验证准确率可能只有 20% 甚至更低,这不是模型坏了,是还在 warmup。
我建议每个 epoch 结束都跑一次验证,并且用 EMA 权重跑,避免训练权重最后的抖动影响判断。
@torch.no_grad() def evaluate(model, loader): model.eval() correct = 0 total = 0 for x, y in loader: x, y = x.to(device), y.to(device) pred = model(x).argmax(dim=1) correct += (pred == y).sum().item() total += y.size(0) return correct / total val_acc = evaluate(ema_model.module, val_loader) print(f"epoch={epoch} val_acc={val_acc:.4f}")这里有个细节:验证时一定要model.eval(),否则 BN 统计量还在更新,验证结果会偏低。FastVIT 的 RepMixer 在训练态也有 BN,不清算这个状态,验证准确率可能上下波动两三个点。
判断训练是否完成,不要只看最终 val_acc。我会记录每个 epoch 之后验证集上每个类别的 recall,如果某一个类别的准确率始终很低,说明数据里可能缺少代表性样本,或者类别间视觉特征太难区分,这时候加训练轮数意义不大,要回去看数据。
最后提醒一点:FastVIT 微调不是轮数越多越好。我跑过一个数据集,40 轮之后验证准确率开始缓慢下降,训练准确率还在涨,这是典型的过拟合信号。遇到这种情况,保存最佳 epoch 的权重,而不是最后一轮权重。这也是 EMA 权重更好用的原因,它的峰值准确率通常比训练权重更稳定。
5. FastVIT 训练避坑与常见问题排查:我先替你踩过这四个坑
5.1 分类头维度不匹配:看似能跑,验证集却在打转
现象:加载 ImageNet 预训练权重后,训练脚本没有报错,但验证准确率一直停留在接近随机水平,比如 10 分类只到 12% 左右。训练 loss 虽然在下降,但速度很慢,像在从头学。
原因:最常见的是你把官方 ImageNet 权重用load_state_dict硬加载到自定义类别的模型里。state_dict会把分类头权重也带上,而自定义模型的最后线性层输出维度和 ImageNet 的 1000 完全不同。PyTorch 对缺 key 或者多 key 会报错,但如果两个模型的网络结构一样,只是分类头输出维度不同,load_state_dict会在 head 处抛 mismatch。
解决:用timm.create_model(..., pretrained=True, num_classes=NUM_CLASSES)加载,timm 会帮你重置分类头。如果你手动加载权重,必须先用strict=False跳过不匹配的 head,再单独初始化新的分类头。加载完一定要打印模型结构确认 head 输出维度,不要相信参数数量差不多就一定是同一个结构。
5.2 混合精度下 Loss 变 NaN:先从学习率入手
现象:开了 AMP 之后,训练到第二个 epoch,loss 突然变成nan,后续所有验证指标都跟着失效。有时候是 loss 直接变 inf,有时候是 optimizer 更新后权重变成 nan。
原因:我遇到过的第一诱因是学习率过高。ViT 系列对学习率比 CNN 敏感,FastVIT 混合了卷积和注意力,注意力部分在高学习率下更容易震荡。第二个诱因是 warmup 太短,模型刚开始还在找方向,就直接给了一个大步长,loss 爆掉。第三个原因比较少见于 FastVIT,但要注意分类头如果是随机初始化,前期梯度会被它带偏。
解决:先把学习率从 1e-3 降到 5e-4 或 3e-4,warmup 从 5 个 epoch 加到 10 个。AMP 相关代码一定要用GradScaler,不能用裸的torch.amp.autocast。此外,EMA 的权重请放在和 model 相同的 device 上,如果 EMA 在 CPU、模型在 GPU,更新时可能出现设备拷贝的隐性 bug,表现为训练正常但验证时精度忽高忽低。
这里有一个可以当后悔药的办法:每 100 步打印一次 loss 和当前学习率,如果 loss 从 2.3 直接跳到 34 或 nan,马上停止训练,不要把整个 epoch 跑完再去看。
5.3 小样本过拟合:森林图像分类里的增强组合
现象:训练集准确率接近 99%,验证集准确率只有 70% 多,两者差距越来越大。你加大数据增强,发现验证集反而掉了,整个训练过程开始有点玄学。
原因:FastVIT 虽然带了卷积归纳偏置,但它仍然有自注意力模块,模型容量对小数据集来说足够大。森林图像分类这种任务,背景中的树叶、光线、阴影都可能被模型当成类别特征,造成过拟合。盲目加 RandAugment 不一定有效,因为auto_augment的强度过高会把原本有判别力的纹理破坏掉。
解决:我一般会分三步走。第一步,在create_transform里把auto_augment="randaug"换成"rand-m9-mstd0.5-inc1"或者直接关掉,保留RandomResizedCrop、翻转和re_prob=0.25的 RandomErasing。第二步,给模型设置一个较小的drop_path_rate,FastVIT 这类混合模型通常默认带 stochastic depth,drop_path_rate=0.1对一千张左右的小数据集很有帮助。第三步,如果数据量低于两千张,冻结前两个 stage 的参数,只微调后面的 block 和分类头,收敛更快也更稳。
注意:冻结 stem 不是所有项目都适用。如果目标图像和 ImageNet 差异很大,比如医学影像或卫星图,反而是全量微调更好。森林图像分类和 ImageNet 比较接近,冻结 stem 通常问题不大。
5.4 导出前后延迟差距大:漏了结构重参数化
现象:PyTorch 里推理一帧只要 10 毫秒,转成 ONNX 后变成 25 毫秒,模型文件也明显变大。更诡异的是,你在导出前已经调用了model.eval(),但 ONNX 图里还是看到很多分支和 BatchNorm 节点。
原因:model.eval()不会把 RepMixer 和训练态卷积里的多分支结构融合掉。FastVIT 的结构重参数化需要显式执行,把训练态的多分支折算成单分支卷积。如果你跳过这一步,导出得到的 ONNX 每一步仍然有分支相加、BN、卷积旁路,算子数量翻倍,延迟自然下不来。
解决:训练完以后,先去 FastVIT 官方仓库里找 reparameterization 相关的流程,按仓库提供的调用方式把训练态模型转换成推理态。转换完以后用torch.script或torch.onnx.export导出。我习惯在导出前加一句校验:随机生成一批同样的输入,在转换前和转换后各跑一次输出,确认最大误差在 1e-4 以内再继续,以免权重融合出问题。
这个坑是 FastVIT 和其他图像分类模型最不一样的地方。ResNet 不需要这个步骤,但 FastVIT 必须做,否则你验收时看到的推理速度会严重误导你。把结构重参数化加进部署流程里,不要只当它是训练完的附加操作。
6. 推理加速进阶:批处理、半精度与 ONNX 导出
FastVIT 训练完成后,我一般不会直接在 PyTorch 里上线,而是先做三件事:批量推理、半精度、ONNX 导出。这三件事互相独立,但都围绕同一个目标:让图像分类模型在真实服务环境里的延迟稳定可控。
批量推理很简单,但很多人会在 validation 代码里一条一条推理,浪费 GPU。把图片按 batch 打包,一次 forward 多个样本,吞吐量会明显改善。用半精度时,只需要把输入images.half(),模型里的权重先转 float16,注意 PyTorch 里要同时保持 model 和 input 都是 half,否则会出现类型不匹配。
import torch @torch.no_grad() def batch_predict(model, loader, half=False): model.eval() results = [] for images, _ in loader: if half: images = images.half() logits = model(images) preds = logits.argmax(dim=1) results.append(preds.cpu()) return torch.cat(results)ONNX 导出是我上线前的最后一道工序。导出前一定先做结构重参数化,再导出,否则 ONNX 里会带着一堆训练态分支。导出时把 batch 维度设成动态,这样训练时用 batch 64,上线时一帧一帧推理也不会报维度错。
model = model.to("cpu").eval() # 这里先执行官方仓库提供的 reparameterize 转换 # 转换后模型应该只剩单分支卷积,无额外 BN 节点 dummy = torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy, "fastvit_image_classifier.onnx", input_names=["images"], output_names=["logits"], dynamic_axes={"images": {0: "batch"}, "logits": {0: "batch"}}, opset_version=17, )导出以后,我建议不要只看 ONNX Runtime 能不能跑通,还要做输出一致性比对。随机抽 50 张验证集图片,分别用 PyTorch 和 ONNX Runtime 推理,比较两者 top-1 预测结果,不一致的数量应该为 0。ONNX 和图优化里的算子融合偶尔会引入微小误差,多数情况不会影响 argmax,但做一次比对能避免线上事故。
我自己的习惯是,每次迭代留一个固定的评估脚本,把类别顺序、模型输入尺寸、normalize 参数全部固化下来。最后一次训练出了个不错的结果,但我换了一个类别顺序重新打包,val acc 居然没变化,上線后才发现预测全错。后来所有项目里我都会在训练完额外保存一份class_to_idx.json,推理服务启动时强制读取它,而不是靠硬编码。希望这次 FastVIT 的实战笔记能帮你在图像分类上少走几段弯路,尤其是那些导出前最容易踩的重参数化和标签顺序问题。
本文还有配套的精品资源,点击获取