☰
从零用PyTorch复现Vision Transformer:图像切块、位置编码与自注意力解析
2026/9/26 4:34:53 网站建设 项目流程

如果你最近在学深度学习,肯定绕不开一个名字: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 x

3.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 x

4. 环境准备与前置条件

复现 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_sizepatch 越小序列越长,计算量增长明显
depth层数增加,参数和显存线性增长
num_heads影响注意力计算方式和模型表达能力
混合精度开启后显存和速度都有改善

10. 常见问题与排查方法

问题现象可能原因排查方式解决方案
安装 PyTorch 后torch.cuda.is_available()返回 FalsePyTorch 装成了 CPU 版运行pip list查看 torch 版本到 PyTorch 官网按 CUDA 版本重装
训练时 OOMbatch_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 视觉主流的后续演化路线,也是面试和论文复现里最常被问到的方向。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询