ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

ViT图像分类实战:PyTorch源码解析与迁移学习训练指南

2026/9/28 16:56:25 拓冰建站 浏览量
ViT图像分类实战:PyTorch源码解析与迁移学习训练指南 简介基于Vision TransformerViT的图像分类实战项目提供完整Python源码与配套数据集主要面向计算机相关专业计科、信息安全、数据科学与大数据、人工智能、通信、物联网等的在校学生、教师及企业开发者既能用于毕业设计、课程设计和大作业也适合作为ViT入门与二次开发的基础。压缩包共15个文件、约31KB代码结构清晰模型定义模块实现ViT网络结构数据集处理模块负责图像读取与预处理训练和预测脚本覆盖完整分类流程辅助工具支持训练与计算量统计另有Markdown说明文档、类别索引JSON及训练过程输出目录便于对照验证。目前已有264人学习/下载代码经作者验证可稳定运行适合新手快速上手。读者可从中学到ViT模型构建、图像分类训练与预测的完整流程也可在此基础上针对自己的数据DIY扩展功能。使用前注意将解压后的项目重命名为英文路径避免因中文路径导致解析错误。1. 用ViT做图像分类这份源码到底解决了什么拿到标题里的 python实现基于ViT的图像分类任务源码数据集可作毕设,运行简单.zip 时你真正需要的是一个能直接跑通、能出指标、能在答辩时讲清楚原理的完整工程而不只是某段孤零零的模型代码。ViTVision Transformer是当下图像分类的主流技术路线之一它把整张图切成固定大小的 patch再像处理句子一样用 Transformer 去建模图像全局关系——这个思路和传统 CNN 的局部卷积完全不同。这套源码最适合两类人一类是毕设选题落在图像分类方向的学生需要快速复现一份能运行、能讲清楚的基线工程另一类是想从 CNN 转向 Transformer 做分类任务的工程师想看看 ViT 的代码组织、训练流程和参数坑。接下来的内容围绕这套方案落地的完整路径展开ViT 的核心结构、数据集组织、训练与评估脚本、必调参数、常见翻车点最后给几个能直接提升完成度的技巧。2. ViT 核心结构拆解图像如何变成 token 序列2.1 为什么 ViT 能成为主流的图像分类路线之一ViT 能火核心原因是它的建模方式和 CNN 有本质区别。CNN 靠卷积核在局部窗口滑动堆很多层才能扩大感受野ViT 在输入端就把整张图切成多个 patch任意两个 patch 之间通过自注意力直接建立联系第一层就能看到全局。这个结构上的差异让 ViT 特别擅长捕捉长距离依赖和物体部件之间的关联而这恰恰是图像分类中区分相似类别时最需要的能力。从工程角度看ViT 的结构和 NLP 里的 BERT 高度同构一套 Transformer 底座同时处理文本和图像也让它成为多模态方向的事实标配。在毕设场景里选 ViT 不需要跟 CNN 比谁在某个小数据集上高两个点——它结构新、概念独立、论文支撑完整答辩时有天然的讲述脉络。源码里最容易让新手困惑的两个点是 patch embedding 和 CLS token把这两个机制吃透整个模型就算入门了。2.2 PatchEmbed 源码实现用卷积一步完成切块与投影一份能跑通的 ViT 源码通常包含四个核心模块PatchEmbed、位置编码、Transformer Encoder 和分类头。先看 PatchEmbed它的任务是把 (B, C, H, W) 的图像变成 (B, N, D) 的 token 序列import torch import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, embed_dim768): super().__init__() self.img_size img_size self.patch_size patch_size self.num_patches (img_size // patch_size) ** 2 # 用 stridepatch_size 的卷积同时完成切块和线性投影 self.proj nn.Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): B, C, H, W x.shape x self.proj(x) # (B, embed_dim, H/p, W/p) x x.flatten(2) # (B, embed_dim, N) x x.transpose(1, 2) # (B, N, embed_dim) return x这里的实现技巧是用kernel_sizepatch_size, stridepatch_size的卷积替代手动切块切块的同时完成线性映射一步到位。以 224×224 输入、patch_size16 为例图像被切成 14×14196 个 patch每个 patch 映射成 768 维向量输出就是 (B, 196, 768)。选择卷积而不是纯 reshape 的原因有两个一是卷积带可学习权重等于在切块时同时做了像素到 embedding 的映射二是 GPU 上卷积的并行效率远高于逐 patch 循环。num_patches这个值会直接影响后面的位置编码维度改输入分辨率或 patch_size 时最容易在这里出维度 mismatch。2.3 CLS token 与位置编码ViT 里的两个标志性设计PatchEmbed 输出 token 序列后ViT 会拼上一个特殊的可学习向量——CLS token然后加上位置编码一起送进 Transformer Encoderclass ViT(nn.Module): def __init__(self, embed_dim768, depth12, num_heads12, num_classes10, img_size224, patch_size16): super().__init__() self.patch_embed PatchEmbed(img_size, patch_size, 3, embed_dim) self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed nn.Parameter( torch.zeros(1, 1 self.patch_embed.num_patches, embed_dim)) self.pos_drop nn.Dropout(p0.1) encoder_layer nn.TransformerEncoderLayer( d_modelembed_dim, nheadnum_heads, dim_feedforwardembed_dim * 4, dropout0.1, activationgelu, batch_firstTrue) self.encoder nn.TransformerEncoder(encoder_layer, num_layersdepth) self.norm nn.LayerNorm(embed_dim) self.head nn.Linear(embed_dim, num_classes) def forward(self, x): B x.shape[0] x self.patch_embed(x) # (B, N, D) cls_tokens self.cls_token.expand(B, -1, -1) x torch.cat([cls_tokens, x], dim1) # (B, N1, D) x self.pos_drop(x self.pos_embed) x self.encoder(x) x self.norm(x) cls_out x[:, 0] # 只看CLS位置 return self.head(cls_out)CLS token 的作用是充当一个可学习的汇总位。它不携带任何具体 patch 的信息但经过 12 层自注意力后它通过 attention 机制看过所有 patch最终状态可以被理解成整张图的全局特征表示再接一层线性分类头输出类别 logits。这个设计和 BERT 里的 [CLS] 完全一致。位置编码用的是可学习参数而非固定的三角函数初始化为 0训练时自己学会每个 patch 的绝对位置关系。注意batch_firstTrue这个参数PyTorch 的 TransformerEncoderLayer 默认输入是 (seq_len, batch, dim)不加这个参数喂 (B, N, D) 会直接维度报错。2.4 数据流动从 224×224 图像到 10 类概率分布把完整前向串联起来图像进入 PatchEmbed变成 (B, 196, 768)拼上 CLS token 变成 (B, 197, 768)加上位置编码后进入 12 层 Transformer Encoder每层内部做 LayerNorm、多头自注意力、MLP 和残差连接。最后一层的输出经过 final norm取 CLS token 对应的那一行过 Linear 头得到 (B, num_classes) 的 logits。如果是推理阶段再接 softmax 就是概率分布。三个细节值得记住。一是dim_feedforward默认是 embed_dim 的 4 倍768 对应 3072这层 MLP 是参数量和显存消耗的最大头。二是 final norm 不能省论文里在 Encoder 后、分类头前加了一层 LayerNorm不加它分类效果会有可感知的下降。三是源码里图像进入模型前一定做过 Normalize用 ImageNet 的 mean[0.485, 0.456, 0.406] 和 std[0.229, 0.224, 0.225]如果数据没做标准化模型前几层的输出会直接爆掉。3. 跑通完整训练流程数据集、训练脚本与评估方法3.1 数据集准备从任意目录结构到 ImageFolder源码必须搭配能跑的数据集才能体现运行简单。图像分类最通用的组织方式是 torchvision 的 ImageFolder 格式根目录下每个类一个文件夹文件夹里放该类所有图片。假设做一个 10 类分类任务data/ ├── train/ │ ├── cat/ # 文件夹名即类别名 │ │ ├── cat_001.jpg │ │ └── cat_002.jpg │ ├── dog/ │ └── ... └── val/ ├── cat/ ├── dog/ └── ...ImageFolder 会自动把文件夹名映射成整数标签不需要手动维护一个类别到 ID 的映射表这是它成为分类任务默认格式的原因。如果原始数据集不是这种结构——比如所有图片都在一个大目录里另配一个 txt 标注文件——需要先做一次转换import os import shutil from pathlib import Path # 把 图片相对路径 类别名 的txt转换为按类别分目录的结构 root Path(raw_images) out Path(data/train) os.makedirs(out, exist_okTrue) with open(labels.txt, r) as f: for line in f: img_path, label line.strip().split() dst_dir out / label os.makedirs(dst_dir, exist_okTrue) shutil.copy(root / img_path, dst_dir / img_path)这段转换脚本的核心逻辑是逐行解析 txt拆出图片路径和类别名再复制到对应类别的目录下。用os.makedirs(dst_dir, exist_okTrue)是为了在批量处理时反复创建已存在的目录不报错。选择复制而不是移动是为了保留原始数据的完整性避免转换脚本出问题时原始数据也被破坏。转换完成后建议随便抽一个类目看一眼图片数量有些原始数据集会存在类别分布极不均匀的情况比如 cat 有 1000 张dog 只有 80 张这种情况后面训练时需要考虑类别均衡。3.2 训练主循环数据加载、loss 计算与参数更新训练脚本是源码里运行简单的核心体现。用 PyTorch 训练 ViT 的标准结构如下import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms # 224是ViT的标准输入尺寸RandomResizedCrop保证尺度多样性 train_tf transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) train_ds datasets.ImageFolder(data/train, transformtrain_tf) train_loader DataLoader(train_ds, batch_size64, shuffleTrue, num_workers4, pin_memoryTrue) model ViT(num_classeslen(train_ds.classes)) model model.cuda() criterion nn.CrossEntropyLoss() optimizer torch.optim.AdamW(model.parameters(), lr3e-4, weight_decay0.05) for epoch in range(epochs): model.train() # 训练模式dropout生效 running_loss 0.0 for images, labels in train_loader: images, labels images.cuda(), labels.cuda() logits model(images) loss criterion(logits, labels) optimizer.zero_grad() loss.backward() optimizer.step() running_loss loss.item() * images.size(0) print(fepoch {epoch 1}, loss: {running_loss / len(train_ds):.4f})训练里最容易忽略的是model.train()必须在每个 epoch 开始前调用。ViT 内部有 dropouttrain 模式下随机失活eval 模式下不启用漏了这句dropout 全程不生效模型在验证集上会显得泛化能力很差。pin_memoryTrue在数据量不大、CPU 算力够时可以让 GPU 拷贝更快它对训练速度的改善比较直观建议保留。CrossEntropyLoss 自带 softmax 操作模型最后一层输出 logits 直接喂进去就行不需要再手动过 softmax这也是新手常有的困惑。3.3 模型结构判断timm 加载还是手写 ViT源码的模型实现一般有两种路线一种像前面手写适合学习和改结构另一种用timm库加载现成的 ViT适合快速出指标。如果你是做毕设而不是发论文强烈建议至少看一眼手写实现因为答辩时老师大概率会问CLS token 是怎么加的位置编码是学习来的还是固定的这些只有读过源码才答得上来。timm 路线同样值得保留一行代码就能加载在 ImageNet 上预训练好的权重在小数据集上的效果比从零训练好得多import timm # 加载在ImageNet上预训练的ViT-Base把分类头换成自己的类别数 model timm.create_model(vit_base_patch16_224, pretrainedTrue, num_classes10)pretrainedTrue会下载在 ImageNet-1k 上训练好的权重num_classes10会自动替换最后的分类头。这条路线最大的优势是省时间——在几百张图的小数据集上从零训练一个 ViT-base 几乎不可能收敛到可用精度但用预训练权重微调随便跑二三十个 epoch 就能出不错的结果。当然手写实现时同样可以加载预训练权重用load_state_dict把 backbone 部分载入再随机初始化分类头。3.4 验证与评估准确率、混淆矩阵和模型保存训练不能只看 loss分类任务的硬指标是验证集准确率。验证脚本需要把模型切到 eval 模式并关闭梯度计算否则 PyTorch 会为验证过程也构建计算图白白占掉显存from sklearn.metrics import confusion_matrix def evaluate(model, val_loader, device): model.eval() correct 0 total 0 all_preds [] all_labels [] with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) logits model(images) preds logits.argmax(dim1) correct (preds labels).sum().item() total labels.size(0) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) acc correct / total cm confusion_matrix(all_labels, all_preds) return acc, cm验证脚本的关键是torch.no_grad()它告诉 PyTorch 不需要为这些计算保存梯度信息能省下大量显存。argmax(dim1)在类别维度上取最大值索引注意dim1而不是dim0——dim0 是 batch 维度。混淆矩阵的价值在于能直观看出模型具体在哪两类之间混淆比如猫和狗互认错这比单一准确率数字更有说服力放进毕设论文里也是很好的分析材料。模型保存用torch.save(model.state_dict(), vit_best.pth)同时保存一份 val_acc 最高的 checkpoint而不是最后一个 epoch 的这个小习惯能帮你留住最佳模型。4. 参数怎么设5 个让 ViT 收敛更稳的必调参数4.1 学习率与 warmupTransformer 最敏感的两个开关ViT 对学习率极其敏感这是 Transformer 系列的共性也是它和 CNN 最明显的训练差异。CNN 用 0.1 的 SGD 学习率经常能跑但 ViT 用 SGD 基本很难收敛主流做法是 AdamW 配 3e-4 左右的学习率再叠加 warmup 策略。warmup 的意思是训练最初的几个 epoch 让学习率从 0 线性爬到设定值。这样做的原因在于ViT 在初始化阶段位置编码和 attention 权重还很粗糙一上来就用大学习率会把参数推到不理想的区域后面很难拉回来。常见的设置是 warmup 占总训练 epoch 的 5% 到 10%。import math def get_lr(epoch, warmup_epochs, total_epochs, base_lr3e-4): if epoch warmup_epochs: return base_lr * (epoch 1) / warmup_epochs # cosine退火后期学习率平滑下降收敛更稳 progress (epoch - warmup_epochs) / (total_epochs - warmup_epochs) return 0.5 * base_lr * (1 math.cos(math.pi * progress))warmup 用线性上升之后用 cosine 退火这是 ViT 训练里最稳定也最常见的组合。cosine 退火比固定学习率好在训练后期不会一直在 loss 曲面边缘震荡能多榨出几个点的准确率。实际调试时如果第一个 epoch loss 就冲到很大先检查是不是 warmup 没接上这是最常见的翻车点之一。4.2 patch size 与输入分辨率计算量的双重杠杆ViT 的计算复杂度跟序列长度 N 的平方成正比而 N (H / patch_size)²。patch_size 设 16、输入 224 时得到 196 个 tokenpatch_size 改成 8同样输入会得到 784 个 token计算量直接翻 4 倍。所以 patch size 是 ViT 里性价比最高的调节旋钮大 patch 跑得快但丢失细节小 patch 精度更高但显存消耗显著上涨。分类任务的一般经验值是 16如果目标物体尺寸小、依赖细节区分可以降到 8 试一轮。输入分辨率同理。224 是 ViT 预训练权重和位置编码的标准尺寸自己从零训练时改大改小都行但用了预训练权重分辨率变了位置编码就得做插值否则维度对不上。有数据表明用 384 分辨率微调 ImageNet 预训练模型能比 224 高 1-2 个点代价是显存翻倍。毕设场景如果只想快速出指标先稳定用 224把精力放在数据增强和训练时长上更划算。改这两个参数时最需要注意的是它们共同决定num_patches进而决定pos_embed的维度改小后加载旧 checkpoint 一定会报 shape mismatch。4.3 depth 与 num_heads网络的深度和宽度怎么配ViT-Base 的默认配置是 12 层、12 头、embed_dim768参数量在 8600 万左右。depth 控制 Transformer Encoder 的层数num_heads 控制多头注意力里把特征分成几个子空间。head 数必须能整除 embed_dim768 除以 12 等于每个头分到 64 维这是原论文的标准配置。增加 depth 能提升表达能力但训练时间和显存线性上涨增加 head 数对精度的提升不明显反而可能在小数据集上加剧过拟合。一个实用的调试顺序是固定 num_heads12把 depth 从 8、12、16 各跑一遍看验证集准确率的上限。如果 12 层和 16 层只差 1 个点以内果断选 12 层理由很实际训练时间省 30%答辩时不会因为训练时间过长被老师质疑实验做得不够。对于毕设级别的数据量ViT-Base 已经属于偏大的配置再往深堆收益很小。4.4 dropout 与 weight decay小数据集上的防过拟合组合ViT 在没有预训练的情况下在小数据集上过拟合非常快典型表现是训练 loss 一直降、验证 loss 回升。两个最有效的控制手段是 dropout 和 weight decay。源码里一般有两个 dropout 位置pos_drop作用于位置编码之后通常设 0.1TransformerEncoderLayer 内部还有 dropout作用于 attention 输出和 MLP 输出通常设 0.1 到 0.3。小数据集建议从 0.1 起步如果验证集过拟合明显加到 0.2 再试。weight decay 在 AdamW 优化器里一般设 0.05这比 CNN 常用的 1e-4 高出不少。原因是 AdamW 把 weight decay 从梯度更新里拆出来单独应用同样的数值在 AdamW 里实际惩罚力度比 SGD 温和所以需要调大到 0.05 级别。完整参数参考如下参数推荐值调节方向learning rate3e-4发散就降一半收敛太慢试 1e-3warmup epochs总 epoch 的 5%-10%小数据集取上限patch size16细节任务降到 8depth / num_heads12 / 12数据量小降到 8 层dropout0.1过拟合时加到 0.2-0.3weight decay0.05过拟合时加到 0.1batch size64显存不够用 amp 或降到 32调参顺序一般先定学习率和 warmup再调 dropout 和 weight decay。这两个组合决定了模型在小数据集上能不能稳定收敛。如果验证集一直上不去优先检查数据和增强策略而不是盲目加大模型。参数调整就是一个反复实验的过程建议每次只动一个变量把训练日志记清楚不然最后根本不知道哪个改动起了作用。5. 常见问题排查从 NaN 到显存爆炸的 5 个真实场景5.1 训练 loss 直接变成 NaN现象第一个 epoch 还没跑完 loss 就变成 NaN或者训练到中途突然炸掉。原因常见有三种。一是学习率过大ViT 对学习率极其敏感AdamW 配 1e-3 以上的学习率在小数据集上很容易发散二是混合精度没有配 GradScalerfp16 在反向传播时梯度下溢成 0 或上溢成 NaN三是数据里有损坏的图片或标签越界比如灰度图被当成 3 通道读入导致维度错误或某个样本的标签超出类别范围。解决把学习率降到 3e-4 或更低重试如果用torch.autocast做混合精度补上 GradScaler在 DataLoader 里加异常捕获打印 loss 变成 NaN 前最后一批数据的 shape 和标签范围定位是不是数据问题。还有一种本地定位法跑一个 batch 的前向打印 logits 的数值范围如果全是 inf 或 NaN问题几乎肯定在网络初始化或输入数据这两端。5.2 GPU 显存不够batch size 只能调到 2现象batch_size32 直接 CUDA out of memory调到 8 还是爆最后只能设 2训练慢到怀疑人生。原因ViT 的显存大头在 attention 矩阵。序列长度 197、12 层、12 个头每个头都要维护 (B, 12, 197, 197) 的 attention 分数batch32 时这部分占了绝大部分显存。再加上 MLP 层和 AdamW 的动量缓存显存消耗比同规模 CNN 高出一截这是结构特性不是代码有 bug。解决优先开 AMP 混合精度显存能省近一半其次把 batch size 降到 16 或 8配合梯度累积模拟大 batch如果数据集分辨率高把输入从 224 降到 192 或 160序列长度变化不大但 attention 矩阵是平方关系节省效果明显。梯度累积的标准写法是accum_steps 4 for step, (images, labels) in enumerate(train_loader): loss criterion(model(images), labels) / accum_steps loss.backward() if (step 1) % accum_steps 0: optimizer.step() optimizer.zero_grad()累积 4 个 step 再更新一次参数等效于 batch size 扩大 4 倍但显存占用不变。注意要把 loss 除以累积步数再做 backward否则等效学习率会偏大。5.3 数据集太小验证集精度还不如随机猜测现象一共几百张图训练几轮之后训练 loss 很低但验证集准确率一直在 20%-30% 徘徊跟随机猜差不多。原因ViT 是典型的数据饥饿型模型没有预训练的情况下需要大量数据才能学到通用的视觉特征。小数据集上它很容易把训练集的噪声细节背下来但对验证集没有泛化能力。这是结构特性不是代码没写对。解决最有效的方法是从预训练权重开始微调几百张图也能出不错的效果。其次加强数据增强RandomResizedCrop、颜色抖动、RandomErasing 都能变相扩充数据。实在不行就换小模型ViT-Tiny 加数据增强在小数据集上的表现通常比 ViT-Base 从零裸跑好。如果是毕设选题建议从一开始就定好预训练权重微调的路线省时省力指标还好看。5.4 加载 checkpoint 报 shape mismatch现象训练完保存的模型换一台机器或改了个配置再加载直接报size mismatch for pos_embed。原因位置编码的 shape 和图像分辨率、patch_size 绑定。保存时的模型是 224 分辨率、patch_size16对应 197 个 token加载时如果改了输入尺寸或 patch_sizepos_embed 的第一个维度就对不上。另一种常见情况是类别数变了head 层权重从 10 类换成 100 类自然不匹配。解决加载时用strictFalse然后手动处理不匹配的层ckpt torch.load(vit_10cls.pth) model ViT(num_classes100) missing, unexpected model.load_state_dict(ckpt, strictFalse) print(missing, unexpected) # 看哪些参数没有加载上load_state_dict 返回的 missing 里会列出 head.weight 和 head.bias这两个本来就应该重新初始化。只有位置编码对不上才是真正的问题需要插值或重新训练。如果改过输入分辨率pos_embed 要从 (197, 768) 插值到新的序列长度比如 14×14 变成 16×16用torch.nn.functional.interpolate对 pos_embed 做二维插值再 reshape 回去。5.5 验证集上没问题测试集上效果骤降现象在验证集上调参调到 95%一上测试集掉到 85%怀疑是不是模型训练出了问题。原因测试集只在最后用一次验证集被反复用来调参验证集的信息已经间接泄露到模型和参数选择里了。反复试参数、选择验证集上最好的模型本质上是在拟合验证集。这是最隐蔽的坑尤其在做毕设时很多人一遍遍在验证集上试最后的测试集效果必然打折。解决把原始数据拆成 train、val、test 三份训练期间只看 val所有调参决策做完后冻结一切流程只跑一次 test。如果数据集太小拆不出三份用交叉验证把多次 val 结果取平均。这个习惯越早养成越好。源码里一般给你 train 和 val 的结构test 集要自己从 train 里再分出来这一步千万别省。6. 进阶技巧迁移学习、注意力可视化与模型导出迁移学习在小数据集上是 ViT 项目性价比最高的优化手段。用 timm 加载 ImageNet 预训练权重替换分类头几百张图也能收敛到可用的水平这比从零训练省下大量时间。微调时先冻结 backbone 训几轮分类头再解冻全部参数用小学习率微调效果通常比直接全参微调更稳。注意力可视化是 ViT 特有的解释性工具。取最后一层所有 head 的 attention 矩阵对 CLS token 那一行做平均就能得到每个 patch 对最终分类的贡献权重。将它 reshape 回 14×14 再上采样到原图尺寸叠加在原图上就能直观看到模型聚焦在哪些区域。这张热图放进毕设论文里比写一整段文字解释模型原理都管用它直观证明了模型确实关注到了目标物体的关键区域而不是靠背景作弊。实现时用 register_forward_hook 抓取 attention 输出或者直接改一层 forward 返回中间变量。模型导出是毕设现场演示前必须做的一步。演示电脑往往没有 GPU把模型导出成 TorchScript 可以脱离训练框架直接加载推理model.eval() scripted torch.jit.script(model) scripted.save(vit_scripted.pt) # CPU 推理 model torch.jit.load(vit_scripted.pt, map_locationcpu) model.eval() with torch.no_grad(): logits model(img_tensor) # (1, num_classes)一个血泪教训导出前一定先model.eval()否则 dropout 和 LayerNorm 的运行状态会被固化推理结果和训练时一致但那是错的。TorchScript 导出后可以明显感觉到 CPU 推理速度的提升因为它省去了 Python 层的调度开销。ViT 的坑和收益都很极端不调参、不用预训练权重它可能让你在第一个 epoch 就心态崩溃锚定好学习率、warmup、dropout 这套组合拳再配合迁移学习它在图像分类任务上的表现又能稳稳超过常规 CNN。我走过最大的弯路是一上来就训练大模型后来学会先跑 tiny 版本验证全链路再切正式模型整个节奏都顺了。希望这篇能帮你把那份源码跑通、跑明白少踩几个我已经替你踩过的坑。本文还有配套的精品资源点击获取