Vision Transformer(ViT)项目源码解析:从模型构建到训练测试全流程 本项目基于PyTorch和TorchVision实现采用官方提供的ViT-B/16预训练模型对EMNIST Letters数据集进行迁移学习实现英文字母分类任务。整个模型架构一、环境配置与依赖导入项目开始首先导入训练过程中需要使用的各种库包括 PyTorch、TorchVision、Scikit-learn、Matplotlib 以及 Seaborn 等。其中PyTorch 负责模型构建和训练TorchVision 提供预训练模型和数据集接口Scikit-learn 用于计算模型评价指标而 Matplotlib 和 Seaborn 则负责实验结果的可视化展示。完成依赖导入后程序会自动检测当前运行环境是否支持 CUDA。如果当前设备安装了 NVIDIA GPU 且 CUDA 环境配置正确程序将自动使用 GPU 进行模型训练否则将退回到 CPU 模式运行。这种自动选择设备的方式使同一份代码能够适配不同硬件环境无需人工修改。import torch import torch.nn as nn import torch.optim as optim import torchvision.transforms as transforms import torchvision from torchvision.models import vit_b_16, ViT_B_16_Weights from sklearn.metrics import confusion_matrix, accuracy_score import matplotlib.pyplot as plt import seaborn as sns import numpy as np from torch.utils.data import random_split # Define device device torch.device(cuda if torch.cuda.is_available() else cpu)这一部分虽然代码量不多但决定了整个项目的运行环境也是后续模型训练能够充分利用计算资源的重要基础。二、Vision Transformer 模型初始化完成环境配置后程序开始构建神经网络模型。# 使用 torchvision 的 ViT (无需 Hugging Face) weights ViT_B_16_Weights.DEFAULT model vit_b_16(weightsweights) # 修改分类头为 27 类 model.heads.head nn.Linear(model.heads.head.in_features, 27) model.to(device)项目直接调用 TorchVision 提供的 ViT-B/16 预训练模型并加载官方发布的预训练权重。这些权重是在大规模 ImageNet 数据集上训练得到的已经具备较强的图像特征提取能力因此无需从零开始训练整个网络。由于预训练模型默认输出 1000 个类别而 EMNIST Letters 数据集只包含 27 个类别因此程序将模型最后的分类层替换为新的全连接层使输出维度与当前数据集保持一致。这种方式属于典型的迁移学习策略即保留 Transformer 主干网络仅修改最后的分类器从而减少训练时间同时提高模型在新任务上的收敛速度和识别精度。模型初始化完成后再将整个模型迁移到前面配置好的计算设备上为后续训练做好准备。三、优化器与损失函数配置模型建立完成后需要确定模型训练过程中采用的优化策略。本项目使用 AdamW 作为优化器。相比传统 AdamAdamW 将权重衰减从梯度更新过程中分离出来能够有效缓解 Transformer 模型训练过程中的过拟合问题因此目前已经成为 Vision Transformer 等模型较为常见的优化算法。损失函数采用 CrossEntropyLoss即交叉熵损失函数。由于本项目属于多类别分类任务因此交叉熵能够较好地衡量模型预测结果与真实标签之间的差异并作为反向传播更新网络参数的重要依据。优化器和损失函数共同决定了模型如何学习也是整个训练过程中最核心的两个组成部分。# Define Optimizer Loss Function optimizer optim.AdamW(model.parameters(), lr5e-5) criterion nn.CrossEntropyLoss()四、数据预处理在深度学习项目中数据预处理直接影响模型训练效果。EMNIST 数据集中的图片大小仅为 28×28并且只有一个灰度通道而 ViT-B/16 的输入要求为 224×224 的 RGB 图像。因此在数据读取过程中需要首先将图片尺寸调整到模型要求的输入大小。随后将原本单通道的灰度图转换为三通道图像以满足预训练模型的输入格式要求。接着将图片转换为 PyTorch Tensor并按照 ImageNet 官方提供的均值和标准差进行归一化处理使输入数据分布与模型预训练阶段保持一致。# Data Transforms (使用 torchvision 推荐的 transforms) transform weights.transforms() # 使用 ViT 官方预处理 # 但 EMNIST 是灰度图需要调整 transform transforms.Compose([ transforms.Resize((224, 224)), transforms.Grayscale(num_output_channels3), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])经过这一系列处理后EMNIST 数据便能够直接输入到 Vision Transformer 中进行训练同时最大程度发挥预训练权重的优势。五、数据集加载与划分完成数据预处理后程序开始加载 EMNIST Letters 数据集。项目分别读取训练集和测试集其中训练集进一步按照 8:2 的比例划分为训练集和验证集。训练集用于模型参数学习验证集用于监控模型在训练过程中的泛化能力而测试集则在整个训练结束后用于最终性能评估。为了提高训练效果训练集对应的 DataLoader 开启随机打乱使每一轮训练的数据顺序都不同有助于提升模型泛化能力验证集和测试集则保持固定顺序以保证每次评估结果的一致性和可重复性。# Load EMNIST dataset train_dataset torchvision.datasets.EMNIST(root./data, splitletters, trainTrue, downloadTrue, transformtransform) test_dataset torchvision.datasets.EMNIST(root./data, splitletters, trainFalse, downloadTrue, transformtransform) # Split into train/val train_size int(0.8 * len(train_dataset)) val_size len(train_dataset) - train_size train_dataset, val_dataset random_split(train_dataset, [train_size, val_size]) train_loader torch.utils.data.DataLoader(train_dataset, batch_size64, shuffleTrue) val_loader torch.utils.data.DataLoader(val_dataset, batch_size64, shuffleFalse) test_loader torch.utils.data.DataLoader(test_dataset, batch_size64, shuffleFalse)这种训练集、验证集、测试集三者分离的方式也是当前深度学习项目中最常见的数据组织方式。六、模型评估模块evaluate为了避免重复编写验证集和测试集的评估代码项目将评估过程封装成一个独立函数。当模型进入评估阶段后首先切换到评估模式同时关闭梯度计算以减少显存占用并提高推理速度。随后模型依次读取数据集中的每个 Batch完成前向传播并计算当前批次的损失值和预测结果。整个数据集遍历结束后程序统计平均损失和分类准确率同时保存所有预测标签和真实标签作为后续绘制混淆矩阵的数据来源。# Evaluation function def evaluate(model, data_loader): model.eval() total_loss 0.0 all_preds [] all_labels [] with torch.no_grad(): for images, labels in data_loader: images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) total_loss loss.item() _, predicted torch.max(outputs, 1) all_preds.extend(predicted.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) avg_loss total_loss / len(data_loader) accuracy accuracy_score(all_labels, all_preds) return avg_loss, accuracy, np.array(all_labels), np.array(all_preds)由于验证集和测试集都调用这一函数因此整个项目避免了重复代码提高了代码复用性也使程序结构更加规范。七、混淆矩阵可视化模块为了更加全面地分析模型分类效果项目实现了混淆矩阵绘制功能。程序根据真实标签和预测标签生成混淆矩阵并利用 Seaborn 绘制热力图同时在每个网格中显示对应的预测数量。相比单纯输出 Accuracy混淆矩阵能够更加直观地展示模型在哪些类别上表现较好哪些类别之间容易发生混淆。例如在字符识别任务中一些形状相近的字母往往更容易出现误分类而这些问题都可以通过混淆矩阵快速发现。# Plot confusion matrix def plot_confusion_matrix(y_true, y_pred, classes, titleConfusion Matrix, save_pathNone): cm confusion_matrix(y_true, y_pred) plt.figure(figsize(14, 12)) sns.heatmap(cm, annotTrue, fmtd, cmapBlues, xticklabelsclasses, yticklabelsclasses, annot_kws{size: 6}) plt.xlabel(Predicted) plt.ylabel(True) plt.title(title) plt.tight_layout() if save_path: plt.savefig(save_path, dpi150) print(fConfusion matrix saved to {save_path}) plt.show() class_labels [chr(i) for i in range(ord(a), ord(z) 1)] [unknown]此外程序还支持将混淆矩阵保存为图片方便后续实验记录和论文绘图。八、模型训练模块train_vit整个训练流程被封装在train_vit()函数中也是整个项目最核心的部分。# Training Loop def train_vit(model, epochs10): train_losses, train_accuracies [], [] val_losses, val_accuracies [], [] for epoch in range(epochs): model.train() running_loss 0.0 train_preds, train_labels [], [] 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() running_loss loss.item() _, predicted torch.max(outputs, 1) train_preds.extend(predicted.cpu().numpy()) train_labels.extend(labels.cpu().numpy()) train_loss running_loss / len(train_loader) train_acc accuracy_score(train_labels, train_preds) train_losses.append(train_loss) train_accuracies.append(train_acc) val_loss, val_acc, _, _ evaluate(model, val_loader) val_losses.append(val_loss) val_accuracies.append(val_acc) print(fEpoch [{epoch1}/{epochs}]) print(f Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}) print(f Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.4f}) print(- * 50) # Plot curves plt.figure(figsize(12, 5)) plt.subplot(1, 2, 1) plt.plot(range(1, epochs1), train_losses, labelTrain Loss, markero) plt.plot(range(1, epochs1), val_losses, labelVal Loss, markers) plt.xlabel(Epoch); plt.ylabel(Loss); plt.title(Loss Curve); plt.legend(); plt.grid(True) plt.subplot(1, 2, 2) plt.plot(range(1, epochs1), train_accuracies, labelTrain Acc, markero) plt.plot(range(1, epochs1), val_accuracies, labelVal Acc, markers) plt.xlabel(Epoch); plt.ylabel(Accuracy); plt.title(Accuracy Curve); plt.legend(); plt.grid(True) plt.tight_layout() plt.savefig(training_curves.png, dpi150) plt.show() return train_losses, train_accuracies, val_losses, val_accuracies在每一个 Epoch 中程序首先将模型切换到训练模式然后按照 Batch 为单位依次读取训练数据完成前向传播、损失计算、梯度清零、反向传播以及参数更新等步骤。每完成一个 Epoch程序都会统计当前训练集上的平均损失和准确率并立即调用前面定义的 evaluate() 函数在验证集上计算对应的性能指标。与此同时程序会将训练集和验证集的 Loss 与 Accuracy 保存下来为后续绘制训练曲线提供数据支持。运行结果九、主程序执行流程完成各个功能模块定义后程序开始进入主流程。首先调用训练函数对模型进行多轮训练并实时记录训练过程中的各项指标。训练结束后再调用评估函数在测试集上计算最终的损失值和分类准确率得到模型最终性能。最后程序分别绘制验证集和测试集的混淆矩阵同时保存训练曲线和混淆矩阵图片方便实验分析和结果展示。# Train train_losses, train_accuracies, val_losses, val_accuracies train_vit(model, epochs10) # Test evaluation print(\n * 50) print(Final Evaluation on Test Set) print( * 50) test_loss, test_acc, test_labels, test_preds evaluate(model, test_loader) print(fTest Loss: {test_loss:.4f}) print(fTest Accuracy: {test_acc:.4f}) print( * 50) # Confusion matrices plot_confusion_matrix(test_labels, test_preds, class_labels, titleTest Set Confusion Matrix, save_pathtest_confusion_matrix.png) val_loss, val_acc, val_labels, val_preds evaluate(model, val_loader) plot_confusion_matrix(val_labels, val_preds, class_labels, titleValidation Set Confusion Matrix, save_pathval_confusion_matrix.png)整个主程序没有包含复杂的业务逻辑而是按照“训练→评估→可视化”的流程依次调用前面定义好的各个模块使整个项目结构层次清晰、职责明确。运行结果总结完整代码import torch import torch.nn as nn import torch.optim as optim import torchvision.transforms as transforms import torchvision from torchvision.models import vit_b_16, ViT_B_16_Weights from sklearn.metrics import confusion_matrix, accuracy_score import matplotlib.pyplot as plt import seaborn as sns import numpy as np from torch.utils.data import random_split # Define device device torch.device(cuda if torch.cuda.is_available() else cpu) # 使用 torchvision 的 ViT (无需 Hugging Face) weights ViT_B_16_Weights.DEFAULT model vit_b_16(weightsweights) # 修改分类头为 27 类 model.heads.head nn.Linear(model.heads.head.in_features, 27) model.to(device) # Define Optimizer Loss Function optimizer optim.AdamW(model.parameters(), lr5e-5) criterion nn.CrossEntropyLoss() # Data Transforms (使用 torchvision 推荐的 transforms) transform weights.transforms() # 使用 ViT 官方预处理 # 但 EMNIST 是灰度图需要调整 transform transforms.Compose([ transforms.Resize((224, 224)), transforms.Grayscale(num_output_channels3), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) # Load EMNIST dataset train_dataset torchvision.datasets.EMNIST(root./data, splitletters, trainTrue, downloadTrue, transformtransform) test_dataset torchvision.datasets.EMNIST(root./data, splitletters, trainFalse, downloadTrue, transformtransform) # Split into train/val train_size int(0.8 * len(train_dataset)) val_size len(train_dataset) - train_size train_dataset, val_dataset random_split(train_dataset, [train_size, val_size]) train_loader torch.utils.data.DataLoader(train_dataset, batch_size64, shuffleTrue) val_loader torch.utils.data.DataLoader(val_dataset, batch_size64, shuffleFalse) test_loader torch.utils.data.DataLoader(test_dataset, batch_size64, shuffleFalse) # Evaluation function def evaluate(model, data_loader): model.eval() total_loss 0.0 all_preds [] all_labels [] with torch.no_grad(): for images, labels in data_loader: images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) total_loss loss.item() _, predicted torch.max(outputs, 1) all_preds.extend(predicted.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) avg_loss total_loss / len(data_loader) accuracy accuracy_score(all_labels, all_preds) return avg_loss, accuracy, np.array(all_labels), np.array(all_preds) # Plot confusion matrix def plot_confusion_matrix(y_true, y_pred, classes, titleConfusion Matrix, save_pathNone): cm confusion_matrix(y_true, y_pred) plt.figure(figsize(14, 12)) sns.heatmap(cm, annotTrue, fmtd, cmapBlues, xticklabelsclasses, yticklabelsclasses, annot_kws{size: 6}) plt.xlabel(Predicted) plt.ylabel(True) plt.title(title) plt.tight_layout() if save_path: plt.savefig(save_path, dpi150) print(fConfusion matrix saved to {save_path}) plt.show() class_labels [chr(i) for i in range(ord(a), ord(z) 1)] [unknown] # Training Loop def train_vit(model, epochs10): train_losses, train_accuracies [], [] val_losses, val_accuracies [], [] for epoch in range(epochs): model.train() running_loss 0.0 train_preds, train_labels [], [] 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() running_loss loss.item() _, predicted torch.max(outputs, 1) train_preds.extend(predicted.cpu().numpy()) train_labels.extend(labels.cpu().numpy()) train_loss running_loss / len(train_loader) train_acc accuracy_score(train_labels, train_preds) train_losses.append(train_loss) train_accuracies.append(train_acc) val_loss, val_acc, _, _ evaluate(model, val_loader) val_losses.append(val_loss) val_accuracies.append(val_acc) print(fEpoch [{epoch1}/{epochs}]) print(f Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}) print(f Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.4f}) print(- * 50) # Plot curves plt.figure(figsize(12, 5)) plt.subplot(1, 2, 1) plt.plot(range(1, epochs1), train_losses, labelTrain Loss, markero) plt.plot(range(1, epochs1), val_losses, labelVal Loss, markers) plt.xlabel(Epoch); plt.ylabel(Loss); plt.title(Loss Curve); plt.legend(); plt.grid(True) plt.subplot(1, 2, 2) plt.plot(range(1, epochs1), train_accuracies, labelTrain Acc, markero) plt.plot(range(1, epochs1), val_accuracies, labelVal Acc, markers) plt.xlabel(Epoch); plt.ylabel(Accuracy); plt.title(Accuracy Curve); plt.legend(); plt.grid(True) plt.tight_layout() plt.savefig(training_curves.png, dpi150) plt.show() return train_losses, train_accuracies, val_losses, val_accuracies # Train train_losses, train_accuracies, val_losses, val_accuracies train_vit(model, epochs10) # Test evaluation print(\n * 50) print(Final Evaluation on Test Set) print( * 50) test_loss, test_acc, test_labels, test_preds evaluate(model, test_loader) print(fTest Loss: {test_loss:.4f}) print(fTest Accuracy: {test_acc:.4f}) print( * 50) # Confusion matrices plot_confusion_matrix(test_labels, test_preds, class_labels, titleTest Set Confusion Matrix, save_pathtest_confusion_matrix.png) val_loss, val_acc, val_labels, val_preds evaluate(model, val_loader) plot_confusion_matrix(val_labels, val_preds, class_labels, titleValidation Set Confusion Matrix, save_pathval_confusion_matrix.png)从整体架构来看本项目遵循了典型的深度学习工程设计思想将模型构建、数据预处理、数据加载、模型训练、模型评估以及结果可视化等功能进行了模块化封装各模块之间相互独立、职责明确具有良好的可读性和可扩展性。这种设计不仅方便后续替换不同的数据集或网络模型也便于增加新的功能模块。例如如果需要将 Vision Transformer 替换为 ResNet、Swin Transformer 或 ConvNeXt只需要修改模型初始化部分即可如果需要引入新的评价指标或训练策略也可以在对应模块中进行扩展而不会影响整个项目的整体结构。对于初学者来说这种模块化组织方式既符合深度学习项目开发的基本规范也能够帮助快速理解一个完整图像分类项目从数据准备到模型训练再到结果分析的完整流程。