如果你最近在学深度学习,肯定绕不开一个名字:Vision Transformer,简称 ViT。网上讲 ViT 的文章很多,但大部分不是堆公式就是只丢一段代码,跑起来之后依然说不清它和 CNN 到底差在哪里。
这次我们换一种方式:直接用 PyTorch 从零复现一个 ViT,把图像切块、位置编码、自注意力这几件事拆开讲清楚。文章会覆盖算法原理、完整源码、训练验证、显存占用、接口调用和常见报错排查,你跟着走一遍,就能在自己的数据集上跑通 ViT 训练流程。
先说结论:ViT 并不是什么高不可攀的模型,它的核心思想一句话能讲完——把图片切成 patch,然后当作 NLP 里的 token 序列,交给标准 Transformer Encoder 处理。真正的难点在于理解 shape 怎么流动。转置卷积的逆思维、Embedding 的拼接逻辑、Multi-Head Attention 的分头计算,这些值一旦对上,整个模型就通了。
1. ViT 核心能力速览
| 能力项 | 说明 |
|---|---|
| 模型全称 | Vision Transformer,由 Google 团队在 2020 年提出 |
| 核心思想 | 将图像划分为固定大小的 patch,展开为序列后输入 Transformer Encoder |
| 典型参数 | ViT-Base 约 8600 万参数,12 层 Encoder,12 头注意力,隐藏维度 768 |
| 推荐硬件 | 建议 6G 以上显存的 NVIDIA GPU;纯 CPU 可跑推理,训练速度会很慢 |
| 支持框架 | PyTorch / TensorFlow / JAX,本文使用 PyTorch |
| 启动方式 | Python 脚本训练 + 推理,可封装为本地 API 服务 |
| 是否支持 API | 可自行封装,Flask / FastAPI 均可 |
| 是否支持批量任务 | 支持,可在 DataLoader 中设置 batch_size |
| 适合场景 | 图像分类、特征提取、迁移学习预训练、下游任务 Backbone |
| 输入规格 | 默认 224x224 RGB 图像,patch 大小 16x16,序列长度 196 |
| 输出内容 | 分类 logits 或特征向量,取决于是否保留 classification head |
如果你之前只跑过 ResNet 这类 CNN 模型,ViT 的显存占用会比同参数量的 CNN 略高,原因在自注意力的计算图更重。本文会专门讲怎么观察和降低显存占用。
2. 适用场景与使用边界
ViT 适合三种人。
第一种是刚入门深度学习、想理解 Transformer 结构如何迁移到视觉领域的同学。ViT 的代码量很小,去掉注释大概 200 行,比任何 CNN 变体都适合读源码。
第二种是需要在自定义图像数据集上做分类的研究者或工程师。相比 CNN,ViT 在大规模数据上更容易获得更高的上限,而且不需要手工设计卷积核。
第三种是做多模态模型的人。CLIP、ALIGN、Flamingo 这些模型的视觉编码器大量采用 ViT 结构,你把 ViT 复现一遍,再看这些模型的源码就不会懵了。
使用边界也很清楚:
- 小数据集上训练 ViT 容易欠拟合,通常需要预训练权重或数据增强策略辅助。
- ViT 对图像分辨率敏感,patch size 和输入尺寸直接决定序列长度,改分辨率会影响参数量和显存。
- 训练过程中涉及数据集读取、模型权重保存等内容,使用的数据集和图片素材必须来自合法渠道,注意版权和隐私合规。
3. ViT 算法原理:从图像到序列
ViT 的前向流程可以拆成五个环节。
3.1 图像切块:Patch Embedding
假设输入图像是 224×224×3,我们设置 patch size 为 16。那么图像会被切成 (224/16)×(224/16)=14×14=196 个 patch,每个 patch 的形状是 16×16×3。
把这 196 个 patch 拉平成一维向量,每个向量长度是 768。这里的 768 就是 ViT-Base 的隐藏维度,也叫 embedding dimension。
用 PyTorch 实现时,通常直接用 Conv2d 完成切块和投影:
class PatchEmbed(nn.Module): """图像切块 + 线性投影""" def __init__(self, in_channels=3, embed_dim=768, patch_size=16): super().__init__() self.patch_size = patch_size self.proj = nn.Conv2d( in_channels, embed_dim, kernel_size=patch_size, stride=patch_size ) def forward(self, x): # x: [B, 3, 224, 224] x = self.proj(x) # [B, 768, 14, 14] x = x.flatten(2) # [B, 768, 196] x = x.transpose(1, 2) # [B, 196, 768] return x用 Conv2d 做 Patch Embedding 是当前最常见实现方式,因为卷积操作天然支持 patch 的切分和投影,速度比手动遍历 patch 快很多。
3.2 添加 CLS Token
Transformer Encoder 处理的是序列,但图像分类最终要输出一个类别标签。ViT 的做法是:在序列最前面拼接一个特殊的 token,叫 CLS token。
这个 token 一开始是随机初始化的,维度也是 768。它和 196 个 patch token 一起送入 Transformer Encoder,经过多层编码后,CLS token 对应位置的输出向量就作为整张图像的全局表征,再接一个全连接层做分类。
class ViT(nn.Module): def __init__(self): super().__init__() self.cls_token = nn.Parameter(torch.zeros(1, 1, 768)) def forward(self, x): # x: [B, 3, 224, 224] x = self.patch_embed(x) # [B, 196, 768] cls_tokens = self.cls_token.expand(x.shape[0], -1, -1) # [B, 1, 768] x = torch.cat([cls_tokens, x], dim=1) # [B, 197, 768] return x3.3 位置编码
序列本身没有顺序概念,patch 被切出来之后,如果不加位置信息,“第 1 个 patch”和“第 5 个 patch”在模型眼里没有区别。ViT 选择给每个 token 加上一个可学习的位置编码向量。
位置编码和 token 向量直接相加,维度保持一致。
self.pos_embed = nn.Parameter(torch.zeros(1, 197, 768)) # 初始化方式参考 DeiT: # nn.init.trunc_normal_(self.pos_embed, std=0.02)这里的 197 = 196 个 patch token + 1 个 CLS token。
3.4 Transformer Encoder Block
每个 Encoder Block 包含 LayerNorm、Multi-Head Self-Attention、MLP 和残差连接。
自注意力是核心。每个 token 生成 Query、Key、Value 三个向量,然后计算 Query 和 Key 的相似度得到注意力权重,加权求和得到输出。多头注意力就是把 768 维分成 12 个 64 维的子空间,每个子空间独立计算注意力,最后拼接起来。
class Attention(nn.Module): """多头自注意力""" def __init__(self, dim=768, num_heads=12): super().__init__() self.num_heads = num_heads self.head_dim = dim // num_heads self.qkv = nn.Linear(dim, dim * 3) self.proj = nn.Linear(dim, dim) def forward(self, x): B, N, C = x.shape qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim) qkv = qkv.permute(2, 0, 3, 1, 4) # [3, B, heads, N, head_dim] q, k, v = qkv[0], qkv[1], qkv[2] attn = (q @ k.transpose(-2, -1)) * (self.head_dim ** -0.5) attn = attn.softmax(dim=-1) x = (attn @ v).transpose(1, 2).reshape(B, N, C) x = self.proj(x) return x这里有一个关键细节:缩放因子head_dim ** -0.5是防止点积结果过大导致 softmax 梯度消失。实现时如果不加,模型通常也能跑通,但收敛速度会差一些,建议保留。
3.5 MLP 与残差连接
每个 Encoder Block 的最后一层是 MLP,包含两个全连接层,中间使用 GELU 激活函数。MLP 的隐藏维度一般是 embedding dim 的 4 倍,ViT-Base 就是 768×4=3072。
随后加上残差连接,保证深层网络的梯度可以流通。
class Mlp(nn.Module): def __init__(self, in_dim=768, hidden_dim=3072): super().__init__() self.fc1 = nn.Linear(in_dim, hidden_dim) self.fc2 = nn.Linear(hidden_dim, in_dim) def forward(self, x): x = self.fc1(x) x = nn.functional.gelu(x) x = self.fc2(x) return x4. 环境准备与前置条件
复现 ViT 之前,先检查环境。
4.1 操作系统与 Python
Ubuntu 20.04 / 22.04、Windows 10/11、macOS 都可以运行。Python 建议使用 3.9 或 3.10,PyTorch 对 3.11 的支持现在也比较稳定,但 3.9/3.10 踩坑最少。
4.2 PyTorch 与 CUDA
如果使用 NVIDIA 显卡,先确认驱动支持 CUDA 版本,再安装对应版本的 PyTorch。判断方式:
python -c "import torch; print(torch.cuda.is_available())"结果为True表示 GPU 可用。为False则检查驱动、CUDA 版本以及 PyTorch 是否安装为 GPU 版。
CPU 机器也可以运行完整代码,只是训练一个 epoch 的时间会长很多,建议先用小数据量验证代码逻辑,再决定是否上 GPU。
4.3 数据集与目录结构
本文以 CIFAR-10 为例,你也可以换成自己的图片数据集。数据集的下载和引用建议使用公开标准数据集,确保来源合法。
建议目录结构如下:
vit-tutorial/ ├── data/ # 数据集存放 ├── checkpoints/ # 模型权重 ├── outputs/ # 推理结果 ├── vit_model.py # 模型定义 ├── train.py # 训练脚本 └── inference.py # 推理脚本5. 完整源码复现:ViT 模型定义
下面是完整的 ViT 模型代码,可以直接保存为vit_model.py。结构上按照 Patch Embedding、CLS Token、位置编码、Encoder Layer、最终分类头依次组合。
import torch import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, in_channels=3, embed_dim=768, patch_size=16): super().__init__() self.patch_size = patch_size self.proj = nn.Conv2d( in_channels, embed_dim, kernel_size=patch_size, stride=patch_size ) def forward(self, x): # x: [B, 3, H, W] x = self.proj(x) # [B, embed_dim, H/p, W/p] x = x.flatten(2) # [B, embed_dim, num_patches] x = x.transpose(1, 2) # [B, num_patches, embed_dim] return x class Attention(nn.Module): def __init__(self, dim=768, num_heads=12): super().__init__() self.num_heads = num_heads self.head_dim = dim // num_heads self.qkv = nn.Linear(dim, dim * 3) self.proj = nn.Linear(dim, dim) def forward(self, x): B, N, C = x.shape qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim) qkv = qkv.permute(2, 0, 3, 1, 4) q, k, v = qkv[0], qkv[1], qkv[2] attn = (q @ k.transpose(-2, -1)) * (self.head_dim ** -0.5) attn = attn.softmax(dim=-1) x = (attn @ v).transpose(1, 2).reshape(B, N, C) x = self.proj(x) return x class Mlp(nn.Module): def __init__(self, dim=768, hidden_dim=3072): super().__init__() self.fc1 = nn.Linear(dim, hidden_dim) self.fc2 = nn.Linear(hidden_dim, dim) def forward(self, x): x = self.fc1(x) x = nn.functional.gelu(x) x = self.fc2(x) return x class EncoderBlock(nn.Module): def __init__(self, dim=768, num_heads=12, mlp_ratio=4.0): super().__init__() self.norm1 = nn.LayerNorm(dim) self.attn = Attention(dim, num_heads) self.norm2 = nn.LayerNorm(dim) self.mlp = Mlp(dim, int(dim * mlp_ratio)) def forward(self, x): x = x + self.attn(self.norm1(x)) x = x + self.mlp(self.norm2(x)) return x class VisionTransformer(nn.Module): def __init__( self, img_size=224, patch_size=16, in_channels=3, num_classes=10, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4.0 ): super().__init__() self.patch_embed = PatchEmbed(in_channels, embed_dim, patch_size) num_patches = (img_size // patch_size) ** 2 self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim)) self.blocks = nn.Sequential(*[ EncoderBlock(embed_dim, num_heads, mlp_ratio) for _ in range(depth) ]) self.norm = nn.LayerNorm(embed_dim) self.head = nn.Linear(embed_dim, num_classes) nn.init.trunc_normal_(self.pos_embed, std=0.02) nn.init.trunc_normal_(self.cls_token, std=0.02) def forward(self, x): B = x.shape[0] x = self.patch_embed(x) # [B, num_patches, embed_dim] cls_tokens = self.cls_token.expand(B, -1, -1) x = torch.cat([cls_tokens, x], dim=1) # [B, num_patches + 1, embed_dim] x = x + self.pos_embed x = self.blocks(x) x = self.norm(x) cls_final = x[:, 0] logits = self.head(cls_final) return logits注意,上面的num_patches是按方形输入计算的。如果输入不是正方形,需要改成(H // patch_size) * (W // patch_size)。实际开发中建议把img_size固定为 224,避免序列长度变化导致位置编码维度不匹配。
6. 训练脚本:从随机初始化开始训练
模型定义好之后,我们需要一份训练脚本。以下代码以 CIFAR-10 为例,包含数据加载、模型初始化、训练循环和 checkpoint 保存。
import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms from vit_model import VisionTransformer def main(): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") transform_train = transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) transform_test = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) train_dataset = datasets.CIFAR10( root="./data", train=True, download=True, transform=transform_train ) test_dataset = datasets.CIFAR10( root="./data", train=False, download=True, transform=transform_test ) train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4, pin_memory=True) test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False, num_workers=4, pin_memory=True) model = VisionTransformer( img_size=224, patch_size=16, in_channels=3, num_classes=10, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4.0 ).to(device) criterion = nn.CrossEntropyLoss() optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0.05) epochs = 30 for epoch in range(epochs): model.train() total_loss = 0.0 correct = 0 total = 0 for images, labels in train_loader: images, labels = images.to(device), labels.to(device) optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() total_loss += loss.item() _, predicted = outputs.max(1) total += labels.size(0) correct += predicted.eq(labels).sum().item() train_acc = 100.0 * correct / total model.eval() test_correct = 0 test_total = 0 with torch.no_grad(): for images, labels in test_loader: images, labels = images.to(device), labels.to(device) outputs = model(images) _, predicted = outputs.max(1) test_total += labels.size(0) test_correct += predicted.eq(labels).sum().item() test_acc = 100.0 * test_correct / test_total print(f"Epoch [{epoch+1}/{epochs}] " f"Loss: {total_loss/len(train_loader):.4f} " f"Train Acc: {train_acc:.2f}% " f"Test Acc: {test_acc:.2f}%") torch.save(model.state_dict(), f"checkpoints/vit_base_epoch_{epoch+1}.pth") print("训练完成") if __name__ == "__main__": main()训练时重点关注两个指标:
Train Acc持续上升,说明模型在学习,没有出现梯度爆炸或网络结构错误。Test Acc在 ViT 随机初始化 + 小数据集的情况下可能偏低,这是正常现象。ViT 在 ImageNet 规模的数据集上才能发挥最大优势,CIFAR-10 上直接随机初始化训练更多是为了验证代码流程。
如果只想验证代码是否跑通,建议用depth=4、embed_dim=192这样的小型配置,或者先用min(1000, len(train_dataset))做一次截断数据测试。
7. 推理验证:加载权重并预测单张图片
训练完成后,写一个推理脚本来验证模型输出。以下脚本会读取一张图片,输出十个类别的 logits,并打印预测类别。
import torch from PIL import Image from torchvision import transforms from vit_model import VisionTransformer checkpoint_path = "checkpoints/vit_base_epoch_30.pth" image_path = "test_images/cat.jpg" class_names = ["airplane", "automobile", "bird", "cat", "deer", "dog", "frog", "horse", "ship", "truck"] transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = VisionTransformer( img_size=224, patch_size=16, in_channels=3, num_classes=10 ).to(device) model.load_state_dict(torch.load(checkpoint_path, map_location=device)) model.eval() image = Image.open(image_path).convert("RGB") x = transform(image).unsqueeze(0).to(device) with torch.no_grad(): logits = model(x) probs = torch.softmax(logits, dim=1) top1 = probs.argmax(dim=1).item() print(f"预测类别: {class_names[top1]}") print(f"置信度: {probs[0][top1].item():.4f}")推理时有两个常见问题:
第一,训练时用了RandomHorizontalFlip等数据增强,推理时不要使用这些随机变换,只保留 Resize、ToTensor、Normalize 即可。
第二,load_state_dict报 shape 不匹配,通常是因为模型参数和 checkpoint 不一致。检查保存时的模型配置是否和加载时的完全一致,包括num_classes、depth、embed_dim。
8. 接口 API 与批量任务
训练好的 ViT 模型可以封装成 HTTP API,方便接到其他系统里。这里用 Flask 给一个最小示例,如果你更习惯 FastAPI,可以直接替换。
import torch from flask import Flask, jsonify, request from PIL import Image from torchvision import transforms from vit_model import VisionTransformer app = Flask(__name__) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = VisionTransformer(img_size=224, patch_size=16, num_classes=10) model.load_state_dict(torch.load("checkpoints/vit_base_epoch_30.pth", map_location=device)) model.to(device) model.eval() transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) class_names = ["airplane", "automobile", "bird", "cat", "deer", "dog", "frog", "horse", "ship", "truck"] @app.route("/predict", methods=["POST"]) def predict(): if "file" not in request.files: return jsonify({"error": "no file uploaded"}), 400 file = request.files["file"] image = Image.open(file.stream).convert("RGB") x = transform(image).unsqueeze(0).to(device) with torch.no_grad(): logits = model(x) probs = torch.softmax(logits, dim=1) top5 = torch.topk(probs[0], k=min(5, len(class_names))) results = [ { "class": class_names[idx], "probability": round(prob.item(), 4) } for idx, prob in zip(top5.indices.tolist(), top5.values.tolist()) ] return jsonify({"predictions": results}) if __name__ == "__main__": app.run(host="0.0.0.0", port=5000)启动服务:
python flask_api.py测试接口:
curl -X POST -F "file=@test_images/cat.jpg" http://127.0.0.1:5000/predict返回示例:
{ "predictions": [ {"class": "cat", "probability": 0.8231}, {"class": "dog", "probability": 0.1042}, {"class": "deer", "probability": 0.0321} ] }批量推理同样可以直接复用 DataLoader。只要把图片目录整理好,使用ImageFolder加载,设置一个较大的batch_size,然后把循环里的单图输入换成 batch 输入即可。批量任务建议在目录里保留输入图片和输出日志,方便失败后重试。
接口服务在生产环境要注意:不要默认绑定0.0.0.0暴露到公网。如果只是本地调试,建议改成127.0.0.1;需要远程访问时,前面加一层网关鉴权,不要让模型接口裸奔。
9. 资源占用与性能观察
ViT 的资源占用,可以从三个维度观察。
9.1 参数量和计算量
ViT-Base(224×224 输入、patch_size=16、depth=12、embed_dim=768、12 heads)的参数量约为 8600 万。对应单张图片的前向计算量按公开实现约为 17.6 GFLOPs 左右,实际数值取决于具体实现和输入分辨率。
作为对比,ResNet-50 参数量约 2500 万,计算量约 4.1 GFLOPs。也就是说,ViT-Base 比 ResNet-50 重不少,这是它在中小数据集上不容易训练好的一个直接原因。
9.2 显存占用观察方法
训练时,可以使用nvidia-smi实时观察显存占用。更精确的方式是在训练脚本里打印模型参数量和激活值显存:
def count_parameters(model): return sum(p.numel() for p in model.parameters() if p.requires_grad) print(f"参数量: {count_parameters(model) / 1e6:.2f}M")显存占用主要由三部分组成:
- 模型参数。
- 优化器状态。
- 前向过程的激活值。
在 batch_size=32、输入为 224×224、ViT-Base 配置下,显存占用通常需要按实际环境测试。如果 OOM,优先调整以下项目:
第一,降低 batch_size。32 改 8,显存基本按比例下降。
第二,使用torch.cuda.amp.autocast()混合精度训练,显存占用可以进一步降低,同时训练速度也会提升。
第三,把num_workers调低,避免数据加载线程抢占内存。
9.3 CPU 推理
CPU 可以运行 ViT 推理,但速度明显下降。以单张 224×224 图片为例,CPU 推理一次通常需要几秒到十几秒不等,具体取决于 CPU 型号和线程数。
GPU 推理则通常在几十毫秒到几百毫秒量级。因此,API 服务部署时建议使用 GPU,纯 CPU 更适合做代码验证和功能调试。
9.4 性能影响因素
影响训练和推理性能的关键参数:
| 参数 | 影响 |
|---|---|
| batch_size | 越大显存占用越高,吞吐量先升后降 |
| 输入分辨率 | 分辨率提高 n 倍,patch 数量提高 n² 倍,注意力计算随之增加 |
| patch_size | patch 越小序列越长,计算量增长明显 |
| depth | 层数增加,参数和显存线性增长 |
| num_heads | 影响注意力计算方式和模型表达能力 |
| 混合精度 | 开启后显存和速度都有改善 |
10. 常见问题与排查方法
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
安装 PyTorch 后torch.cuda.is_available()返回 False | PyTorch 装成了 CPU 版 | 运行pip list查看 torch 版本 | 到 PyTorch 官网按 CUDA 版本重装 |
| 训练时 OOM | batch_size 过大或输入分辨率过高 | 观察nvidia-smi显存占用 | 降低 batch_size、开启 AMP、减小输入尺寸 |
| 加载权重时报 shape 不匹配 | 模型配置和 checkpoint 不一致 | 打印两个 state_dict 的 key 和 shape | 统一模型配置后再保存和加载 |
| 训练时 Loss 不下降 | 学习率设置不合适或数据未归一化 | 检查数据预处理和优化器参数 | 调整学习率、检查 Normalize、缩减模型规模测试 |
| CPU 推理非常慢 | 未使用 GPU 或推理方式不当 | 打印device确认设备 | 换 GPU,或对 CPU 使用torch.set_num_threads |
| API 请求返回 500 | 图片读取失败或模型未加载 | 查看服务日志 | 检查图片格式,确认 checkpoint 路径 |
| 小数据集 Test Acc 很低 | ViT 随机初始化在数据量小时难收敛 | 查看训练集准确率是否同步低 | 换小模型配置,或加载预训练权重做迁移学习 |
| 端口被占用 | 之前有服务未关闭 | lsof -i:5000或netstat -ano | 换端口或结束占用进程 |
单独说一下小数据集的问题。ViT 在小数据集上随机初始化训练,效果大概率不如 ResNet。不是因为代码写错了,而是 Transformer 结构本身需要大量数据学习归纳偏置。工程上常用的解法有三种:
- 使用 ImageNet 预训练权重,然后冻结前几层只训练分类头。
- 增大数据增强强度,Mixup、CutMix、RandAugment 都可以明显提升 ViT 在小数据集上的表现。
- 改用 DeiT 这种带蒸馏机制的 ViT 变体,蒸馏教师模型可以是 CNN,能补齐一部分归纳偏置。
11. 最佳实践与使用建议
第一次复现 ViT,不要直接开大模型。先跑一个微型配置,验证整体流程。
model = VisionTransformer( img_size=224, patch_size=16, num_classes=10, embed_dim=192, depth=4, num_heads=4, mlp_ratio=4.0 )embed_dim=192、depth=4 的微型 ViT 参数量在 500 万级别,训练速度和显存占用都更适合调试。
后续工程化时,建议做这五件事:
第一,目录规范。data、checkpoints、outputs分开存放,模型权重不要和数据集混在一起。
第二,固定随机种子。训练前设置好torch.manual_seed(0),保证每次复现结果一致。
第三,训练日志结构化。把 loss、accuracy、learning rate 写入 TensorBoard 或者直接存 CSV,不要只打印到控制台。
第四,批量任务加日志和失败重试。API 服务里要对异常输入做 try-except,批量推理要给每个文件记录成功或失败状态。
第五,涉及人脸、声音、版权素材时,确认数据和模型的合法授权,不做越权的数据采集和生成。训练用的图片集必须是合法获取且有使用授权的数据。
12. 总结与下一步
ViT 的复现难度不在模型本身,而在于理解图像到序列的转换过程。只要把 Patch Embedding、CLS token、位置编码和 Multi-Head Attention 这四个模块的 shape 变化弄清楚,整个模型就通了。
建议你按下面的顺序验证一遍:
- 先跑通微型 ViT,确认训练 Loss 能下降。
- 再用 ViT-Base 在 CIFAR-10 上完整训练,记录显存和准确率。
- 然后把模型封装成 Flask API,用 curl 测试单张图片预测。
- 最后加上混合精度训练和批量推理,把它接到你自己的数据集和业务场景里。
写代码验证模型逻辑时,注意不要把训练集和测试集的处理混在一起。推理阶段不加载数据增强,只保留 Resize、ToTensor 和 Normalize,这是最容易忽略的坑。
后续如果你想继续深入,可以按这几个方向扩展:一是看 DeiT,了解蒸馏如何弥补 ViT 在小数据集上的不足;二是看 Swin Transformer,了解层级化窗口注意力如何降低计算量;三是看 DINO,了解自监督 ViT 如何提取通用视觉特征。这三条线基本覆盖了目前 ViT 视觉主流的后续演化路线,也是面试和论文复现里最常被问到的方向。