news 2026/7/26 21:12:16

Vision Transformer实战:从零开始搭建ViT图像分类模型(PyTorch版)

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Vision Transformer实战:从零开始搭建ViT图像分类模型(PyTorch版)

Vision Transformer实战:从零搭建PyTorch图像分类模型

当卷积神经网络(CNN)长期统治计算机视觉领域时,Transformer架构的横空出世彻底改变了游戏规则。2020年,Google Research提出的Vision Transformer(ViT)首次证明纯Transformer架构在图像分类任务上可以超越CNN。本文将带您从零开始,用PyTorch实现一个完整的ViT模型,深入解析每个关键组件的实现细节。

1. ViT核心原理与架构设计

ViT的核心思想是将图像处理为序列数据。与传统CNN逐层提取局部特征不同,ViT将图像分割为固定大小的图块(patches),将这些图块线性投影后加上位置编码,送入标准的Transformer编码器进行处理。

关键创新点

  • 图像分块处理:将224×224图像分割为16×16的图块(共196个),每个图块展平为768维向量
  • 可学习分类标记:在序列开头添加特殊[CLS]标记,其最终状态用于分类
  • 位置编码:由于Transformer本身不包含位置信息,必须显式添加位置编码

实验表明,当训练数据足够大时(如JFT-300M),ViT的性能显著优于同规模的CNN模型,尤其在捕捉长距离依赖关系方面表现突出

2. 环境准备与数据预处理

在开始编码前,我们需要配置开发环境并准备数据集:

pip install torch==1.12.0 torchvision==0.13.0 pip install numpy matplotlib tqdm

使用CIFAR-10数据集进行演示,但需要注意原始ViT设计输入尺寸为224×224:

from torchvision import transforms, datasets # 数据增强与归一化 train_transform = transforms.Compose([ transforms.Resize(224), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 加载CIFAR-10数据集 train_set = datasets.CIFAR10(root='./data', train=True, download=True, transform=train_transform) train_loader = torch.utils.data.DataLoader(train_set, batch_size=32, shuffle=True, num_workers=2)

3. ViT核心组件实现

3.1 图块嵌入层

将图像分割为图块并线性投影:

class PatchEmbedding(nn.Module): def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768): super().__init__() self.img_size = img_size self.patch_size = patch_size self.n_patches = (img_size // patch_size) ** 2 self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size) def forward(self, x): # x: [B, C, H, W] x = self.proj(x) # [B, E, H/P, W/P] x = x.flatten(2) # [B, E, N] x = x.transpose(1, 2) # [B, N, E] return x

3.2 位置编码与分类标记

class ViTEmbeddings(nn.Module): def __init__(self, config): super().__init__() self.patch_embeddings = PatchEmbedding( img_size=config.img_size, patch_size=config.patch_size, embed_dim=config.hidden_size ) # 可学习的分类标记 self.cls_token = nn.Parameter(torch.zeros(1, 1, config.hidden_size)) # 位置编码 self.position_embeddings = nn.Parameter( torch.zeros(1, self.patch_embeddings.n_patches + 1, config.hidden_size) ) self.dropout = nn.Dropout(config.dropout_rate) def forward(self, x): batch_size = x.shape[0] # 生成图块嵌入 patch_embeds = self.patch_embeddings(x) # [B, N, E] # 扩展分类标记到batch维度 cls_tokens = self.cls_token.expand(batch_size, -1, -1) # [B, 1, E] # 拼接分类标记 embeddings = torch.cat((cls_tokens, patch_embeds), dim=1) # [B, N+1, E] # 添加位置编码 embeddings += self.position_embeddings return self.dropout(embeddings)

3.3 多头注意力机制实现

class MultiHeadAttention(nn.Module): def __init__(self, config): super().__init__() self.num_heads = config.num_heads self.head_dim = config.hidden_size // config.num_heads self.query = nn.Linear(config.hidden_size, config.hidden_size) self.key = nn.Linear(config.hidden_size, config.hidden_size) self.value = nn.Linear(config.hidden_size, config.hidden_size) self.out = nn.Linear(config.hidden_size, config.hidden_size) self.dropout = nn.Dropout(config.attention_dropout_rate) def forward(self, x): batch_size, seq_len, embed_dim = x.shape # 线性投影得到Q, K, V q = self.query(x) # [B, N, E] k = self.key(x) # [B, N, E] v = self.value(x) # [B, N, E] # 重塑为多头形式 q = q.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2) k = k.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2) v = v.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2) # 计算注意力分数 scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim) attn_weights = torch.softmax(scores, dim=-1) attn_weights = self.dropout(attn_weights) # 应用注意力权重 context = torch.matmul(attn_weights, v) context = context.transpose(1, 2).contiguous().view(batch_size, seq_len, embed_dim) # 最终线性层 output = self.out(context) return output

4. 完整ViT模型集成

将所有组件组合成完整模型:

class ViT(nn.Module): def __init__(self, config): super().__init__() self.config = config # 嵌入层 self.embeddings = ViTEmbeddings(config) # Transformer编码器 encoder_layer = nn.TransformerEncoderLayer( d_model=config.hidden_size, nhead=config.num_heads, dim_feedforward=config.mlp_dim, dropout=config.dropout_rate ) self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=config.num_layers) # 分类头 self.classifier = nn.Linear(config.hidden_size, config.num_classes) def forward(self, x): # 嵌入层 embeddings = self.embeddings(x) # Transformer编码 encoded = self.encoder(embeddings) # 使用分类标记进行分类 cls_token = encoded[:, 0, :] logits = self.classifier(cls_token) return logits

5. 模型训练与优化

配置训练参数并实现训练循环:

# 模型配置 class ViTConfig: img_size = 224 patch_size = 16 hidden_size = 768 num_heads = 12 mlp_dim = 3072 num_layers = 12 dropout_rate = 0.1 attention_dropout_rate = 0.0 num_classes = 10 # 初始化模型 config = ViTConfig() model = ViT(config).to(device) # 定义损失函数和优化器 criterion = nn.CrossEntropyLoss() optimizer = torch.optim.AdamW(model.parameters(), lr=3e-5, weight_decay=0.01) # 训练循环 for epoch in range(10): model.train() running_loss = 0.0 for i, (inputs, labels) in enumerate(train_loader): inputs, labels = inputs.to(device), labels.to(device) optimizer.zero_grad() outputs = model(inputs) loss = criterion(outputs, labels) loss.backward() optimizer.step() running_loss += loss.item() if i % 100 == 99: print(f'Epoch {epoch+1}, Batch {i+1}, Loss: {running_loss/100:.4f}') running_loss = 0.0

6. 关键问题与解决方案

问题1:位置编码如何影响模型性能?

ViT使用可学习的位置编码而非固定的正弦编码。实验表明:

位置编码类型Top-1准确率
可学习编码78.5%
正弦编码77.9%
无位置编码65.3%

问题2:如何处理不同尺寸的输入图像?

可以通过调整图块大小或使用金字塔结构:

# 动态调整图块大小 def adjust_patch_size(img_size, target_patches=196): patch_size = int(math.sqrt(img_size[0]*img_size[1]/target_patches)) return patch_size

问题3:如何加速ViT训练?

  • 使用混合精度训练
  • 采用梯度检查点技术
  • 使用更大的batch size配合学习率warmup
# 混合精度训练示例 scaler = torch.cuda.amp.GradScaler() with torch.cuda.amp.autocast(): outputs = model(inputs) loss = criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()

7. 模型评估与可视化

实现注意力可视化可以帮助理解模型工作原理:

def visualize_attention(image, model, layer_idx=0, head_idx=0): # 注册钩子获取注意力权重 attention_weights = [] def hook(module, input, output): attention_weights.append(output[1]) handle = model.encoder.layers[layer_idx].self_attn.register_forward_hook(hook) # 前向传播 with torch.no_grad(): _ = model(image.unsqueeze(0).to(device)) handle.remove() # 可视化特定头的注意力 attn = attention_weights[0][0, head_idx, 0, 1:].reshape(14, 14) plt.imshow(attn.cpu().numpy(), cmap='hot') plt.colorbar() plt.show()

8. 进阶优化技巧

技巧1:知识蒸馏

使用大型ViT模型(如ViT-Large)作为教师模型,蒸馏到小型ViT:

# 定义蒸馏损失 def distillation_loss(student_logits, teacher_logits, temp=2.0): soft_teacher = F.softmax(teacher_logits/temp, dim=-1) soft_student = F.log_softmax(student_logits/temp, dim=-1) return F.kl_div(soft_student, soft_teacher, reduction='batchmean') * (temp**2)

技巧2:混合架构

结合CNN和Transformer的优势:

class HybridViT(nn.Module): def __init__(self): super().__init__() # 使用CNN提取低级特征 self.cnn_backbone = torchvision.models.resnet18(pretrained=True) self.cnn_backbone = nn.Sequential(*list(self.cnn_backbone.children())[:-2]) # ViT处理高级特征 self.vit = ViT(config) def forward(self, x): cnn_features = self.cnn_backbone(x) b, c, h, w = cnn_features.shape cnn_features = cnn_features.view(b, c, h*w).transpose(1, 2) return self.vit(cnn_features)

技巧3:数据高效训练

  • 使用CutMix数据增强
  • 应用MixUp正则化
  • 采用RandAugment自动增强策略
# CutMix实现示例 def cutmix_data(x, y, alpha=1.0): lam = np.random.beta(alpha, alpha) rand_index = torch.randperm(x.size(0)) target_a = y target_b = y[rand_index] bbx1, bby1, bbx2, bby2 = rand_bbox(x.size(), lam) x[:, :, bbx1:bbx2, bby1:bby2] = x[rand_index, :, bbx1:bbx2, bby1:bby2] lam = 1 - ((bbx2 - bbx1) * (bby2 - bby1) / (x.size()[-1] * x.size()[-2])) return x, target_a, target_b, lam
版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/7/14 14:35:37

极链云服务器跑Python代码保姆级教程:从文件上传到命令行执行

极链云服务器Python代码全流程实战指南:从零部署到高效执行 在当今数据驱动的时代,云服务器已成为开发者不可或缺的工具。极链云以其友好的用户界面和即用型环境配置,为Python开发者提供了快速上手的解决方案。本文将带您深入探索从文件上传到…

作者头像 李华
网站建设 2026/7/14 14:35:42

计算机毕业设计:Python新浪新闻智能采集推荐系统 Django框架 Vue Selenium爬虫 可视化 大数据 数据分析(建议收藏)✅

博主介绍:✌全网粉丝10W,前互联网大厂软件研发、集结硕博英豪成立工作室。专注于计算机相关专业项目实战6年之久,选择我们就是选择放心、选择安心毕业✌ > 🍅想要获取完整文章或者源码,或者代做,拉到文章底部即可与…

作者头像 李华
网站建设 2026/7/14 14:35:41

嵌入式开发三大编译链接问题实战解析

附录:嵌入式开发常见问题与工程实践解决方案在嵌入式固件开发过程中,尤其是基于ARM Cortex-M系列MCU(如STM32F103、STM32F407等)配合Keil MDK、IAR EWARM或GCC工具链进行开发时,开发者常遭遇若干看似琐碎却极具迷惑性的…

作者头像 李华
网站建设 2026/7/14 14:35:55

终极指南:使用Grunt构建LiteGraph.js的完整自动化流程

终极指南:使用Grunt构建LiteGraph.js的完整自动化流程 【免费下载链接】litegraph.js A graph node engine and editor written in Javascript similar to PD or UDK Blueprints, comes with its own editor in HTML5 Canvas2D. The engine can run client side or …

作者头像 李华
网站建设 2026/7/14 14:35:42

Three_Level_NPC_Inverter_基于MATLAB/Simulink的仿真模型

Three_Level_NPC_Inverter:基于MATLAB/Simulink的三电平中性点钳位(NPC)逆变器仿真模型。 仿真条件:MATLAB/Simulink R2015b,购买前如需转成低版本格式请提前告知,谢谢。三电平NPC逆变器在新能源并网、电机…

作者头像 李华
网站建设 2026/7/14 14:35:40

CTF-加密与解密(十七):音乐符号背后的秘密

1. 音乐符号加密:CTF中的另类挑战 第一次在CTF比赛中遇到音乐符号加密题时,我盯着满屏的♪♫♩♬符号发呆了整整十分钟。这种非传统加密方式就像是用钢琴键盘敲出的摩斯密码,既让人摸不着头脑又充满艺术感。音乐符号加密属于文本替换加密的变…

作者头像 李华