news 2026/9/28 5:17:06

ViT图像分类课设实战:从零手写Patch Embedding到ONNX部署

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
ViT图像分类课设实战:从零手写Patch Embedding到ONNX部署

简介:本资源是一套基于Vision Transformer(ViT)的图像分类完整项目实现,面向计算机相关专业在校学生、教师及初入AI领域的从业者,适用于课程设计、毕业设计、大作业等实践场景。项目代码经实测可正常运行,涵盖数据加载、ViT模型构建、训练与预测全流程,并附带配套数据集与详细说明文档,兼顾入门学习与进阶修改需求。压缩包共32个文件,含12个核心Python源码(如vit_model.py、train.py、predict.py)、6个编译缓存文件、4个Markdown说明文档(含项目介绍、使用指南)、2个JSON配置文件(class_indices.json)及辅助文本文件,整体仅66KB,轻量易部署。目前已有325人学习下载,结构清晰、模块解耦良好,特别适合理解Transformer在CV领域的落地逻辑,同时提供FLOPs计算、日志记录、数据预处理等实用工具脚本,便于快速复现与二次开发。

1. Vision Transformer 图像分类项目:为什么课设选它,不是因为“新”,而是因为它真能跑通、真能调、真能讲清楚原理

你手头这个.zip文件——“基于vision transformer图像分类项目python实现源码+数据集(课设新项目).zip”——不是又一个套壳 ResNet 的伪创新。它背后是 Vision Transformer(ViT)在中小规模图像分类任务上的真实落地路径:不用 GPU 集群,一块 RTX 3060 就能训完;不依赖 Hugging Face 全栈黑盒,从 patch embedding 到 class token 拼接,每行代码都可打断点调试;数据集不是 ImageNet-1K,而是你本地解压即用的 3 类花卉或 5 类工业零件图,带 train/val/test 三级目录和标准 label.txt。这不是论文复现,是课设级工程闭环:数据加载 → ViT 模块定义 → 训练循环 → 混淆矩阵可视化 → ONNX 导出 → 单图推理脚本。我带过 17 届本科生做类似课题,翻车最多的是把 ViT 当成“换掉 backbone 的 ResNet”来用——结果 patch size 设错、pos embedding 维度对不上、class token 初始化为零导致梯度消失。这篇笔记就拆开这个 zip 包,告诉你:ViT 在课设场景下,到底哪几行代码不能改、哪些参数必须手调、哪些报错信息一出现就能定位到 loader 还是 model 定义层。适合正在写课程设计报告、需要答辩演示、又不想被问“你这 ViT 和 CNN 本质区别在哪”的同学;也适合想用最小成本验证 ViT 是否适配自己产线质检图像的工程师。


2. 从零构建 ViT 分类器:不调用 transformers 库,手写核心模块与数据流

ViT 不是魔法,它只是把图像切成小块(patch),再用 Transformer 编码器处理这些块序列。课设项目里,避免引入 Hugging Face transformers 库是明智选择——它封装太深,ViTModel.from_pretrained()一行代码背后藏着 200 行初始化逻辑,debug 时根本不知道attention_probs是从哪一层输出的。我们手写,才能控制每个环节。下面三步是整个项目的骨架,全部基于 PyTorch 原生 API,无第三方模型库依赖。

2.1 Patch Embedding 层:图像切块不是简单 reshape,关键在 stride 与 padding 对齐

ViT 的第一步是将输入图像(如 224×224×3)切分为固定大小的 patch(如 16×16),每个 patch 展平为向量,再经线性层映射到 embedding 维度(如 768)。这里最容易被忽略的是patch 切分必须严格整除图像尺寸,否则nn.Unfold或F.unfold会报错或漏采边缘像素。

import torch import torch.nn as nn import torch.nn.functional as F 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 # 必须整除! # 关键:用 Conv2d 实现 patch 切分,比 unfold 更稳定 self.proj = nn.Conv2d( in_chans, embed_dim, kernel_size=patch_size, stride=patch_size # 严格 stride=patch_size,避免重叠或漏采 ) self.norm = nn.LayerNorm(embed_dim) def forward(self, x): # x: [B, C, H, W] -> [B, embed_dim, H//p, W//p] x = self.proj(x) # 自动完成切块 + 线性映射 x = x.flatten(2).transpose(1, 2) # [B, n_patches, embed_dim] x = self.norm(x) return x

参数说明:img_size必须能被patch_size整除(224÷16=14,成立);若你用自定义数据集(如 320×240 工业图),需先 resize 到 224×224 或改patch_size=20(320÷20=16,240÷20=12);embed_dim决定后续 Transformer 层宽度,课设用 768 足够,1024 会显著增加显存压力。

2.2 ViT Encoder Block:复用标准 TransformerEncoderLayer,但 class token 必须手动拼接

PyTorch 的nn.TransformerEncoderLayer已实现 MSA + FFN,我们只需在其输入前插入 class token,并在输出后提取它。class token 不是 learnable parameter,而是可训练的 embedding 向量,且必须在每个 block 输入前 concat。

class ViTEncoderBlock(nn.Module): def __init__(self, embed_dim=768, num_heads=12, mlp_ratio=4.0, dropout=0.1): super().__init__() self.norm1 = nn.LayerNorm(embed_dim) self.attn = nn.MultiheadAttention(embed_dim, num_heads, dropout=dropout, batch_first=True) self.norm2 = nn.LayerNorm(embed_dim) mlp_hidden_dim = int(embed_dim * mlp_ratio) self.mlp = nn.Sequential( nn.Linear(embed_dim, mlp_hidden_dim), nn.GELU(), nn.Dropout(dropout), nn.Linear(mlp_hidden_dim, embed_dim), nn.Dropout(dropout) ) def forward(self, x): # x: [B, n_patches+1, embed_dim], class token 在 index 0 x_norm = self.norm1(x) attn_out, _ = self.attn(x_norm, x_norm, x_norm) # 注意:q,k,v 都是 x_norm x = x + attn_out x = x + self.mlp(self.norm2(x)) return x class VisionTransformer(nn.Module): def __init__(self, img_size=224, patch_size=16, in_chans=3, num_classes=1000, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4.0, dropout=0.1): super().__init__() self.patch_embed = PatchEmbedding(img_size, patch_size, in_chans, embed_dim) self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) # 可学习的 class token self.pos_embed = nn.Parameter(torch.zeros(1, self.patch_embed.n_patches + 1, embed_dim)) self.pos_drop = nn.Dropout(p=dropout) self.blocks = nn.Sequential(*[ ViTEncoderBlock(embed_dim, num_heads, mlp_ratio, dropout) for _ in range(depth) ]) self.norm = nn.LayerNorm(embed_dim) self.head = nn.Linear(embed_dim, num_classes) def forward(self, x): B = x.shape[0] x = self.patch_embed(x) # [B, n_patches, embed_dim] # 拼接 class token: [B, 1, embed_dim] + [B, n_patches, embed_dim] -> [B, n_patches+1, embed_dim] cls_tokens = self.cls_token.expand(B, -1, -1) x = torch.cat((cls_tokens, x), dim=1) x = x + self.pos_embed # 加位置编码 x = self.pos_drop(x) x = self.blocks(x) # [B, n_patches+1, embed_dim] x = self.norm(x) x = x[:, 0] # 取 class token 输出 x = self.head(x) return x

关键逻辑:cls_token.expand(B, -1, -1)确保 batch 维度动态扩展;x[:, 0]提取 class token(索引 0),这是 ViT 分类的核心——所有 patch 信息通过 attention 聚合到这个 token 上;self.pos_embed维度必须是[1, n_patches+1, embed_dim],否则广播失败。

2.3 数据加载与预处理:课设数据集的标准化流程,不是直接套用 torchvision.transforms

课设数据集通常只有几百张图,且类别不平衡(如“缺陷品”仅 30 张,“良品”有 200 张)。直接transforms.RandomHorizontalFlip()可能加剧 imbalance。我们采用分层采样 + 自适应增强:

from torch.utils.data import Dataset, DataLoader, WeightedRandomSampler from torchvision import transforms from PIL import Image import os class CustomImageDataset(Dataset): def __init__(self, root_dir, transform=None, is_train=True): self.root_dir = root_dir self.transform = transform self.is_train = is_train self.classes = sorted(os.listdir(root_dir)) self.class_to_idx = {cls: i for i, cls in enumerate(self.classes)} self.samples = [] self.weights = [] for cls in self.classes: cls_path = os.path.join(root_dir, cls) if not os.path.isdir(cls_path): continue for img_name in os.listdir(cls_path): if img_name.lower().endswith(('.png', '.jpg', '.jpeg')): self.samples.append((os.path.join(cls_path, img_name), self.class_to_idx[cls])) # 训练集按类别反比赋权,缓解 imbalance if is_train: self.weights.append(1.0 / len([s for s in self.samples if s[1] == self.class_to_idx[cls]])) def __len__(self): return len(self.samples) def __getitem__(self, idx): img_path, label = self.samples[idx] image = Image.open(img_path).convert('RGB') if self.transform: image = self.transform(image) return image, label # 课设级预处理:不追求 SOTA,追求稳定收敛 train_transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomAffine(degrees=10, translate=(0.1, 0.1), scale=(0.9, 1.1)), # 轻微形变,比 flip 更鲁棒 transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) # ImageNet 标准化,通用性强 ]) val_transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 构建 dataloader,启用 weighted sampler train_dataset = CustomImageDataset("data/train", transform=train_transform, is_train=True) val_dataset = CustomImageDataset("data/val", transform=val_transform, is_train=False) # 计算每个样本权重(已内置在 dataset 中) sampler = WeightedRandomSampler(train_dataset.weights, len(train_dataset.weights), replacement=True) train_loader = DataLoader(train_dataset, batch_size=16, sampler=sampler, num_workers=2, pin_memory=True) val_loader = DataLoader(val_dataset, batch_size=16, shuffle=False, num_workers=2, pin_memory=True)

为什么不用 RandomHorizontalFlip?在工业缺陷检测中,缺陷方向具有物理意义(如焊缝裂纹沿特定轴向),水平翻转可能生成无效样本;RandomAffine提供更自然的空间扰动。WeightedRandomSampler直接解决课设常见小样本 imbalance,比在 loss 里加 class weight 更早介入。


3. 训练与验证:课设级超参设置、loss 设计与 early stopping 实现

ViT 训练不像 CNN 那样“随便设个 lr=0.01 就能跑”,它的优化器、学习率衰减、warmup 步骤都需针对性调整。课设项目资源有限(单卡、无分布式),必须用最简配置达成收敛。

3.1 ViT 专用优化器:AdamW 替代 SGD,weight decay 必须分离

ViT 对 weight decay 极其敏感。CNN 中常对所有参数统一加 decay,但 ViT 的 LayerNorm 和 bias 不应被正则化,否则训练不稳定。PyTorch 1.12+ 支持no_weight_decay参数,但课设项目建议手动分离参数组:

def get_param_groups(model): """分离可 decay 和不可 decay 的参数""" decay = [] no_decay = [] for name, param in model.named_parameters(): if not param.requires_grad: continue if "bias" in name or "LayerNorm" in name or "ln_" in name: # ln_ 是 ViT 中 LayerNorm 的别名 no_decay.append(param) else: decay.append(param) return [ {'params': decay, 'weight_decay': 0.05}, {'params': no_decay, 'weight_decay': 0.0} ] model = VisionTransformer(num_classes=5) # 假设你的课设是 5 分类 optimizer = torch.optim.AdamW(get_param_groups(model), lr=1e-4, betas=(0.9, 0.999), eps=1e-8)

参数依据:ViT 论文推荐weight_decay=0.05,lr=1e-4(非 1e-3)——因为 ViT 的 embedding 层参数量大,高 lr 易导致 embedding 梯度爆炸;betas=(0.9, 0.999)是 AdamW 默认值,无需改动;eps=1e-8防止除零,课设数据噪声大,保持默认即可。

3.2 学习率 warmup + cosine decay:前 10 个 epoch 线性提升,后 40 个 epoch 余弦衰减

ViT 需要 warmup 让 embedding 层和 position embedding 逐步适应。课设总 epoch 设为 50,warmup 占 20%(10 epoch):

from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR # warmup scheduler: 0 -> base_lr over 10 epochs warmup_scheduler = LinearLR(optimizer, start_factor=1e-3, end_factor=1.0, total_iters=10) # main scheduler: cosine decay from base_lr to 1e-6 over remaining 40 epochs main_scheduler = CosineAnnealingLR(optimizer, T_max=40, eta_min=1e-6) # 组合 scheduler class CombinedScheduler: def __init__(self, warmup, main, warmup_epochs=10): self.warmup = warmup self.main = main self.warmup_epochs = warmup_epochs self.current_epoch = 0 def step(self): self.current_epoch += 1 if self.current_epoch <= self.warmup_epochs: self.warmup.step() else: self.main.step() scheduler = CombinedScheduler(warmup_scheduler, main_scheduler, warmup_epochs=10)

为什么不用 StepLR?StepLR 在固定 epoch 降 lr,ViT 收敛曲线平滑,cosine decay 更匹配其优化轨迹;eta_min=1e-6防止后期 lr 过小导致停滞。

3.3 损失函数与评估指标:Focal Loss 缓解 imbalance,混淆矩阵驱动 debug

课设数据集常存在类别不平衡(如“异常”样本极少),CrossEntropyLoss 会偏向多数类。Focal Loss 通过调节难易样本权重,提升 minority class 召回率:

class FocalLoss(nn.Module): def __init__(self, alpha=1, gamma=2, reduction='mean'): super().__init__() self.alpha = alpha self.gamma = gamma self.reduction = reduction def forward(self, inputs, targets): ce_loss = F.cross_entropy(inputs, targets, reduction='none') pt = torch.exp(-ce_loss) focal_weight = (self.alpha * (1 - pt) ** self.gamma) focal_loss = focal_weight * ce_loss if self.reduction == 'mean': return focal_loss.mean() elif self.reduction == 'sum': return focal_loss.sum() else: return focal_loss criterion = FocalLoss(alpha=1, gamma=2) # gamma=2 是 ViT 论文推荐值

评估不止看 accuracy:课设答辩时老师必问“你的模型在少数类上表现如何?”。训练循环中必须记录 per-class precision/recall/f1:

from sklearn.metrics import confusion_matrix, classification_report import numpy as np def validate(model, val_loader, device): model.eval() all_preds = [] all_labels = [] with torch.no_grad(): for images, labels in val_loader: images, labels = images.to(device), labels.to(device) outputs = model(images) _, preds = torch.max(outputs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) # 计算混淆矩阵 cm = confusion_matrix(all_labels, all_preds) report = classification_report(all_labels, all_preds, target_names=train_dataset.classes, output_dict=True) return cm, report # 在每个 epoch 结束后调用 cm, report = validate(model, val_loader, device) print(f"Epoch {epoch}: Macro F1 = {report['macro avg']['f1-score']:.4f}")

注意:classification_report的output_dict=True返回字典,方便提取report['class 0']['recall']等细粒度指标,答辩时可直接展示“缺陷类召回率 89.2%”。


4. 避坑指南:课设 ViT 项目中 5 个高频翻车点与血泪解决方案

ViT 课设项目翻车,90% 发生在环境配置、数据加载、模型定义三个环节。以下是我在指导 17 届学生时记录的真实报错,按“现象→原因→解决”结构整理,每一条都对应 zip 包里某个文件的修改点。

4.1 现象:RuntimeError: Expected 4-dimensional input for 4-dimensional weight [768, 3, 16, 16], but got 3-dimensional input of size [3, 224, 224] instead

  • 原因:DataLoader返回的images张量维度是[C, H, W](单图),而非[B, C, H, W](batch)。常见于测试脚本中直接torchvision.io.read_image()后未 unsqueeze(0)。
  • 解决:检查inference.py或test_single_image.py中图像加载部分:
    # 错误写法 img = read_image("test.jpg") # 返回 [3, 224, 224] output = model(img) # 模型期待 [1, 3, 224, 224] # 正确写法 img = read_image("test.jpg").unsqueeze(0) # [1, 3, 224, 224] output = model(img)

4.2 现象:训练 loss 不下降,始终在 1.6~1.7 波动(5 分类任务,logit 输出未 softmax)

  • 原因:nn.CrossEntropyLoss内部已包含 softmax + log,若模型输出层额外加了nn.Softmax(),会导致 double softmax,logit 被压缩至 [0,1] 区间,梯度极小。
  • 解决:确认模型forward函数末尾不要加nn.Softmax():
    # 错误 def forward(self, x): x = self.head(x) return F.softmax(x, dim=1) # 删除这一行! # 正确 def forward(self, x): x = self.head(x) return x # CrossEntropyLoss 自动处理

4.3 现象:ValueError: Expected more than 1 value per channel when training, got input size torch.Size([1, 768])来自 LayerNorm

  • 原因:DataLoader的batch_size=1,而LayerNorm在 batch 维度归一化,当 batch=1 时,variance=0导致除零。
  • 解决:课设中batch_size至少设为 4(RTX 3060 显存足够);若必须 batch=1,临时替换LayerNorm为nn.BatchNorm1d(需调整输入维度):
    # 在 PatchEmbedding.__init__ 中 # self.norm = nn.LayerNorm(embed_dim) # 注释掉 self.norm = nn.BatchNorm1d(embed_dim) # 替换为 BN,输入 [B, embed_dim] 时有效 # 在 forward 中 # x = self.norm(x) # 改为 x = self.norm(x.transpose(1, 2)).transpose(1, 2) # BN 要求 [B, C, L]

4.4 现象:CUDA out of memory即使 batch_size=4 也报错

  • 原因:ViT 的 attention 计算复杂度为 O(n²),n 是 patch 数(224÷16=14,n=196,n²=38416)。若embed_dim=1024或num_heads=16,显存暴涨。
  • 解决:课设级务必使用embed_dim=768, num_heads=12, depth=6(非论文 default 的 12 层)。修改VisionTransformer.__init__:
    # 原始(易爆显存) model = VisionTransformer(depth=12, embed_dim=1024) # 课设安全配置 model = VisionTransformer(depth=6, embed_dim=768) # 显存占用降 60%

4.5 现象:验证集 accuracy 95%,但实际预测全错(所有样本被判为同一类)

  • 原因:DataLoader的shuffle=True仅在train_loader中启用,val_loader若也 shuffle,则all_labels和all_preds索引错位,混淆矩阵计算失效。
  • 解决:严格保证val_loader的shuffle=False,并在validate()函数中用enumerate确认顺序:
    # validate 函数开头添加 print("Val loader length:", len(val_loader)) for i, (imgs, lbls) in enumerate(val_loader): print(f"Batch {i}: labels shape {lbls.shape}, first 3 labels {lbls[:3]}") break # 确保输出为 [16, ...] 且标签连续,非随机打乱

5. 模型部署与课设答辩技巧:ONNX 导出、单图推理与可视化解释

课设验收不只是“跑通”,更要让老师看到:你能把 ViT 从训练环境迁移到生产环境,能解释模型为什么这么判,能应对真实场景的输入变化。以下三步是答辩加分项,全部基于 zip 包内代码扩展,无需额外库。

5.1 导出 ONNX 模型:脱离 PyTorch 环境,为后续 C++/Java 部署铺路

ONNX 是跨框架部署的标准格式。ViT 导出需注意 dynamic axes 设置,否则推理时 batch size 固定:

# export_onnx.py import torch import torch.onnx model = VisionTransformer(num_classes=5) model.load_state_dict(torch.load("best_model.pth")) model.eval() # 构造 dummy input: [1, 3, 224, 224] dummy_input = torch.randn(1, 3, 224, 224) # 导出,指定 dynamic batch size torch.onnx.export( model, dummy_input, "vit_classification.onnx", input_names=["input"], output_names=["output"], dynamic_axes={ "input": {0: "batch_size"}, # 第 0 维(batch)可变 "output": {0: "batch_size"} }, opset_version=13 # ViT 需要 opset 12+ ) print("ONNX export success!")

验证 ONNX 模型:用onnxruntime简单测试:

import onnxruntime as ort import numpy as np ort_session = ort.InferenceSession("vit_classification.onnx") outputs = ort_session.run(None, {"input": dummy_input.numpy()}) print("ONNX output shape:", outputs[0].shape) # 应为 [1, 5]

5.2 单图推理脚本:支持命令行传入图片路径,输出 top-3 预测及置信度

答辩演示时,老师会说“现场传一张图看看”。写一个predict.py,一键搞定:

# predict.py import argparse import torch from PIL import Image from torchvision import transforms def main(): parser = argparse.ArgumentParser() parser.add_argument("--model", type=str, default="best_model.pth") parser.add_argument("--image", type=str, required=True) parser.add_argument("--classes", type=str, default="data/classes.txt") # 每行一个类别名 args = parser.parse_args() # 加载模型 model = VisionTransformer(num_classes=5) model.load_state_dict(torch.load(args.model)) model.eval() # 加载类别名 with open(args.classes, "r") as f: classes = [line.strip() for line in f.readlines()] # 预处理 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]) ]) img = Image.open(args.image).convert("RGB") img_tensor = transform(img).unsqueeze(0) # [1, 3, 224, 224] # 推理 with torch.no_grad(): output = model(img_tensor) probs = torch.nn.functional.softmax(output, dim=1)[0] top3_prob, top3_idx = torch.topk(probs, 3) # 输出 print(f"Image: {args.image}") for i in range(3): print(f"Top-{i+1}: {classes[top3_idx[i]]} ({top3_prob[i]:.3f})") if __name__ == "__main__": main()

使用方式:python predict.py --image test_flower.jpg --model best_model.pth,输出清晰可读,答辩时直接终端运行。

5.3 Class Activation Mapping(CAM)可视化:解释 ViT “看哪里做决策”

ViT 没有传统 CNN 的 feature map,但可以用 attention rollout 或 grad-CAM 变体。课设级推荐Attention Rollout:利用最后一层 attention weights,反向传播 patch 重要性:

def attention_rollout(model, img_tensor, head_fusion="mean", discard_ratio=0.9): """简化版 attention rollout,返回每个 patch 的重要性分数""" model.eval() attentions = [] def hook_fn(module, input, output): # output 是 [B, num_heads, seq_len, seq_len],取 mean heads att = output[0].mean(0) # [seq_len, seq_len] attentions.append(att) # 注册 hook 到所有 MultiheadAttention hooks = [] for blk in model.blocks: hooks.append(blk.attn.register_forward_hook(hook_fn)) with torch.no_grad(): _ = model(img_tensor) # 清理 hook for hook in hooks: hook.remove() # rollout: 从最后一层开始,逐层累积 attention result = attentions[-1] # [seq_len, seq_len] for i in range(len(attentions) - 2, -1, -1): result = torch.matmul(result, attentions[i]) # 丢弃最低 90% 的注意力权重,保留 top 10% w, h = 14, 14 # 224//16 mask = result[0, 1:].reshape(w, h).cpu().numpy() # class token 对其他 patch 的 attention mask = (mask - mask.min()) / (mask.max() - mask.min() + 1e-8) mask = np.clip(mask, 0, 1) # 上采样到 224x224 from scipy.ndimage import zoom mask = zoom(mask, (224/w, 224/h), order=1) return mask # 使用示例 img_pil = Image.open("test.jpg").convert("RGB") img_tensor = transform(img_pil).unsqueeze(0) mask = attention_rollout(model, img_tensor) # 可视化 import matplotlib.pyplot as plt plt.figure(figsize=(10, 5)) plt.subplot(1, 2, 1) plt.imshow(img_pil) plt.title("Original") plt.axis("off") plt.subplot(1, 2, 2) plt.imshow(img_pil, alpha=0.5) plt.imshow(mask, cmap='jet', alpha=0.5) plt.title("Attention Rollout") plt.axis("off") plt.show()

答辩话术:“老师,这张图模型认为花瓣边缘区域贡献最大(红色区域),这和我们标注的‘花瓣缺损’缺陷类型一致,说明 ViT 学到了有意义的局部特征,不是靠背景纹理作弊。”


我带课设时发现,学生最怕的不是代码写不出来,而是答辩时被问“你这个 ViT 和 ResNet 有什么本质不同”。后来我要求所有人必须在报告里放一张 attention rollout 图,并手写解释“class token 如何聚合 patch 信息”。ViT 的价值不在参数量,而在它强迫你思考:图像的本质是局部纹理,还是全局关系?这个项目 zip 包里的代码,每一行都在回答这个问题。希望帮到你。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/28 5:16:48

C# WinForms横版卷轴游戏开发实战:GDI+游戏循环与碰撞检测

简介&#xff1a;本资源是一份基于C#开发的横版卷轴动作冒险游戏——《勇士传说》完整源码&#xff0c;面向计算机专业本科生、游戏开发初学者及毕业设计实践者&#xff0c;提供可运行、可调试、可二次开发的实战型学习范例。压缩包共2023个文件&#xff0c;主体为853个Unity a…

作者头像 李华
网站建设 2026/9/28 5:16:21

TCP、UDP、ICMP、HTTP、HTTPS:五个协议的分工与排障思路

先说一个我上个月参与的排查场景&#xff1a;业务方反馈线上接口大面积超时&#xff0c;后端拿着错误日志说“上游返回502”&#xff0c;网络组说交换机端口有轻微丢包&#xff0c;测试的同学补了一句“我这边ping网关延迟有点高”。五个人开了半小时会&#xff0c;结论依然是“…

作者头像 李华
网站建设 2026/9/28 5:15:20

茗茶JSP课设项目全解析:从环境搭建到Tomcat部署实战

茗茶文化网站挂上"JSP"这个技术标签&#xff0c;老JavaWeb人应该一眼就能看穿它的全貌&#xff1a;前台茶叶展示加购物车&#xff0c;后台分类管理加内容维护&#xff0c;数据库用MySQL&#xff0c;跑在Tomcat上&#xff0c;典型的课程设计项目。项目包里那串"q…

作者头像 李华
网站建设 2026/9/28 5:14:50

CPU亲和性设置:解决大小核调度问题,让程序固定跑在大核上

1. 为什么CPU会“大核闲着、小核跑断腿”&#xff1f;1.1 大小核架构的本质&#xff1a;P核与E核到底差在哪先说一个可能很多人没细想的问题&#xff1a;现在市面上主流的“大小核”CPU&#xff0c;到底是怎么个“大”法、“小”法&#xff1f;Intel从12代酷睿开始全面转向混合…

作者头像 李华
网站建设 2026/9/28 5:13:50

Flutter数独App统计卡片组件设计:从Cubit状态管理到OpenHarmony适配实践

最近在把一款数独游戏App往OpenHarmony平台上迁移&#xff0c;顺手把首页那组统计卡片组件重写了一遍。之前这堆卡片其实是临时拼的&#xff0c;Container套着几行Text&#xff0c;数据直接从全局变量里读&#xff0c;看起来能用&#xff0c;但页面一切换就掉状态&#xff0c;数…

作者头像 李华