
1. 项目概述ViTVision Transformer是近年来计算机视觉领域的一项突破性技术它成功将自然语言处理中的Transformer架构迁移到了图像识别任务中。不同于传统的CNN卷积神经网络ViT通过将图像分割成多个patch并线性嵌入然后直接将这些patch序列输入到标准的Transformer编码器中进行处理。这个项目将带你从零开始用最简单的代码实现一个完整的ViT图像分类流程。不同于大多数教程只关注模型构建部分我们会涵盖从环境配置、数据准备、模型训练到结果可视化的全流程。特别适合有以下需求的开发者想快速体验ViT在图像分类任务中的表现需要一套可直接运行的完整代码模板希望理解ViT各组件在实际应用中的配置方法提示本教程使用PyTorch框架所有代码都在Colab上测试通过。即使没有GPU也可以跟着步骤完整跑通。2. 环境搭建与依赖安装2.1 基础环境配置首先确保你的Python版本在3.7以上。推荐使用conda创建虚拟环境conda create -n vit_env python3.8 conda activate vit_env安装核心依赖库pip install torch torchvision torchaudio pip install matplotlib numpy tqdm对于可视化部分我们还会用到pip install opencv-python pillow2.2 ViT模型实现选择虽然可以手动实现ViT但为了快速上手我们使用timm库提供的预实现版本pip install timm这个库包含了多种ViT变体的预训练模型包括ViT-Base (ViT-B/16)ViT-Large (ViT-L/16)ViT-Huge (ViT-H/14)注意如果你在Windows上遇到安装问题可能需要先安装Microsoft C Build Tools。3. 数据准备与预处理3.1 数据集选择为了演示方便我们使用CIFAR-10数据集。虽然ViT通常在大规模数据集上表现更好但CIFAR-10足够展示整个流程from torchvision import datasets, transforms transform transforms.Compose([ transforms.Resize(224), # ViT默认输入尺寸 transforms.ToTensor(), transforms.Normalize(mean[0.5, 0.5, 0.5], std[0.5, 0.5, 0.5]) ]) train_data datasets.CIFAR10(root./data, trainTrue, downloadTrue, transformtransform) test_data datasets.CIFAR10(root./data, trainFalse, downloadTrue, transformtransform)3.2 数据加载器配置创建数据加载器时根据你的GPU内存调整batch_sizefrom torch.utils.data import DataLoader batch_size 64 # 11GB显存可设置为128 train_loader DataLoader(train_data, batch_sizebatch_size, shuffleTrue) test_loader DataLoader(test_data, batch_sizebatch_size)4. ViT模型构建与训练4.1 模型初始化使用timm库加载预定义的ViT模型import timm model timm.create_model(vit_base_patch16_224, pretrainedTrue, num_classes10)关键参数说明vit_base_patch16_224: 使用base规模的ViTpatch大小为16x16输入图像224x224pretrainedTrue: 加载在ImageNet上预训练的权重num_classes10: 我们的CIFAR-10有10个类别4.2 训练配置设置优化器和学习率调度器import torch.optim as optim from torch.optim.lr_scheduler import CosineAnnealingLR optimizer optim.AdamW(model.parameters(), lr1e-4, weight_decay0.01) scheduler CosineAnnealingLR(optimizer, T_max10) criterion torch.nn.CrossEntropyLoss()4.3 训练循环完整的训练epoch实现def train_epoch(model, loader, optimizer, criterion, device): model.train() total_loss 0 for inputs, labels in loader: inputs, labels inputs.to(device), labels.to(device) optimizer.zero_grad() outputs model(inputs) loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() return total_loss / len(loader)5. 可视化与结果分析5.1 注意力可视化ViT最有趣的部分是其注意力机制。我们可以可视化不同head的注意力图import numpy as np import matplotlib.pyplot as plt def visualize_attention(image, model, layer_idx6, head_idx0): # 注册hook获取注意力权重 attentions [] def hook(module, input, output): attentions.append(output[1].detach().cpu()) # 获取注意力权重 handle model.blocks[layer_idx].attn.attn_drop.register_forward_hook(hook) # 前向传播 model.eval() with torch.no_grad(): _ model(image.unsqueeze(0)) handle.remove() # 可视化特定head的注意力 attn attentions[0][0, head_idx, 0, 1:] # 忽略cls token patch_size model.patch_embed.patch_size[0] grid_size int(np.sqrt(attn.shape[0])) attn attn.reshape(grid_size, grid_size) attn torch.nn.functional.interpolate( attn.unsqueeze(0).unsqueeze(0), scale_factorpatch_size, modebilinear )[0, 0] plt.imshow(image.permute(1, 2, 0).cpu().numpy()) plt.imshow(attn, cmaphot, alpha0.5) plt.axis(off)5.2 训练过程监控使用Matplotlib绘制损失和准确率曲线def plot_metrics(train_losses, test_accs): fig, (ax1, ax2) plt.subplots(1, 2, figsize(12, 4)) ax1.plot(train_losses, labelTrain Loss) ax1.set_xlabel(Epoch) ax1.set_ylabel(Loss) ax1.legend() ax2.plot(test_accs, labelTest Accuracy) ax2.set_xlabel(Epoch) ax2.set_ylabel(Accuracy (%)) ax2.legend() plt.tight_layout() plt.show()6. 完整训练流程实现将所有组件整合成一个完整的训练脚本def train_model(model, train_loader, test_loader, epochs10, devicecuda): model model.to(device) optimizer optim.AdamW(model.parameters(), lr1e-4, weight_decay0.01) scheduler CosineAnnealingLR(optimizer, T_maxepochs) criterion torch.nn.CrossEntropyLoss() train_losses [] test_accs [] for epoch in range(epochs): # 训练阶段 model.train() epoch_loss 0 for inputs, labels in tqdm(train_loader, descfEpoch {epoch1}): inputs, labels inputs.to(device), labels.to(device) optimizer.zero_grad() outputs model(inputs) loss criterion(outputs, labels) loss.backward() optimizer.step() epoch_loss loss.item() scheduler.step() train_losses.append(epoch_loss / len(train_loader)) # 测试阶段 model.eval() correct 0 total 0 with torch.no_grad(): for inputs, labels in test_loader: inputs, labels inputs.to(device), labels.to(device) outputs model(inputs) _, predicted torch.max(outputs.data, 1) total labels.size(0) correct (predicted labels).sum().item() accuracy 100 * correct / total test_accs.append(accuracy) print(fEpoch {epoch1}: Loss{train_losses[-1]:.4f}, Accuracy{accuracy:.2f}%) return train_losses, test_accs7. 实际应用技巧与问题排查7.1 显存不足的解决方案如果遇到CUDA out of memory错误可以尝试减小batch_size建议从32开始尝试使用混合精度训练from torch.cuda.amp import autocast, GradScaler scaler GradScaler() # 在训练循环中 with autocast(): outputs model(inputs) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()尝试更小的ViT变体如vit_small_patch16_2247.2 提升准确率的技巧数据增强添加随机裁剪、水平翻转等transform_train transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean[0.5, 0.5, 0.5], std[0.5, 0.5, 0.5]) ])学习率预热前几个epoch逐步提高学习率标签平滑减轻过拟合criterion torch.nn.CrossEntropyLoss(label_smoothing0.1)7.3 常见错误与修复尺寸不匹配错误确保输入图像确实是224x224检查transform中是否有Resize(224)NaN损失降低学习率添加梯度裁剪torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)预测结果随机检查模型是否处于eval模式确保没有忘记model.eval()8. 模型部署与应用8.1 保存与加载模型保存最佳模型torch.save({ model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), }, best_vit.pth)加载模型进行推理checkpoint torch.load(best_vit.pth) model.load_state_dict(checkpoint[model_state_dict]) model.eval()8.2 创建预测API使用Flask创建简单的预测接口from flask import Flask, request, jsonify import io from PIL import Image app Flask(__name__) app.route(/predict, methods[POST]) def predict(): if file not in request.files: return jsonify({error: no file uploaded}) file request.files[file].read() image Image.open(io.BytesIO(file)) # 应用相同的transform image transform(image).unsqueeze(0) with torch.no_grad(): output model(image) _, predicted torch.max(output, 1) return jsonify({class: predicted.item()})8.3 转换为ONNX格式便于跨平台部署dummy_input torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy_input, vit_model.onnx, input_names[input], output_names[output], dynamic_axes{ input: {0: batch_size}, output: {0: batch_size} } )