从零训练ViT做CIFAR10图像分类:patch、位置编码与训练避坑指南
2026/9/24 21:15:34 网站建设 项目流程

简介:基于视觉变换器(ViT)实现CIFAR-10图像分类的训练与验证Python源码包,面向计算机视觉初学者、人工智能方向学生及相关从业者,可用于课程设计、毕业设计或项目初期算法验证。资源核心是一个完整可运行的Python脚本,完整覆盖数据加载、模型搭建、训练循环与验证评估等关键环节,并附带文本说明文件,帮助使用者快速了解代码结构、依赖关系和运行顺序,以及各模块的调用关系。压缩包内共2个文件,以Python脚本为主、txt说明为辅,整个包仅2KB,轻量便捷,适合下载后直接阅读、运行与二次开发。目前已有525人学习/下载,代码经过运行验证,功能稳定,能够帮助读者深入理解视觉Transformer在标准分类任务上的落地流程,也可作为拓展其他视觉任务的修改基底,对计算机视觉入门和毕设答辩均有参考意义,实用价值较高。

1. 为什么是 ViT + CIFAR10:一个 30 分钟能跑通的小型 Transformer 闭环

第一个反直觉的结论是:直接把为 224x224 图像设计的 ViT 预训练权重搬来跑 CIFAR10,十有八九会翻车,远不如从零训练一个 patch 缩小后的 ViT 稳。这套组合解决的是一个小而完整的闭环:用 Vision Transformer 在 32x32 的 CIFAR10 上完成从 patch 切分、位置编码、Transformer 编码器到训练和验证的 python 落地。它的价值在于参数量只有几百万,单卡甚至 CPU 都能跑,半小时内能看到收敛曲线,特别适合第一次接触 ViT 的 pytorch 使用者,也适合熟手快速做轻量 baseline。网上搜 Vit 可能带出电力设备仿真里的同名缩写,注意别下错代码包。下面给的代码按常见做法组织,可直接复现。

2. ViT 结构拆解:patch 大小、位置编码与分类头在 CIFAR10 上的 4 个关键点

2.1 把 32x32 的 CIFAR10 切成 patch:为什么 ViT 原版在这里会失效

ViT 原版的设计基准是 ImageNet:输入 224x224,patch 尺寸 16x16,切完后得到 196 个 token,这个序列长度足够让全局自注意力发挥作用。但 CIFAR10 是 32x32,如果沿用 patch=16,整张图只能切成 2x2 共 4 个 patch,序列长度从 196 暴跌到 4,Transformer 的全局注意力机制基本失去意义,模型退化成一个极弱的局部特征提取器。

常见的做法有两种:一是把 CIFAR10 图像 resize 到 224x224 再喂给预训练 ViT,这样 token 序列恢复正常,但会破坏 CIFAR10 原始的 32x32 分布,放大和插值带来的伪影会让模型学到不该学的噪声,训练成本和显存开销也成倍增加;二是保持图像尺寸不变,把 patch 缩小到 4x4 或 8x8。我一般用 patch=4,这样序列长度是 8x8=64,注意力计算量可控,信息粒度也匹配小图场景。

patch embedding 的经典实现是用一个 stride 等于 kernel_size 的卷积层,卷积核尺寸就是 patch 尺寸,输出通道就是 embedding 维度:

import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, in_channels=3, patch=4, embed_dim=384): super().__init__() self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch, stride=patch) def forward(self, x): # x: [B, 3, 32, 32] x = self.proj(x) # [B, embed_dim, 8, 8] return x.flatten(2).transpose(1, 2) # [B, 64, embed_dim]

把卷积输出的二维特征图打平成序列,再转置成[batch, num_patches, embed_dim],这一步就被称为 patch embedding。这里的逻辑是:卷积的每个输出位置对应原图的一个 patch 区域,通道维记录 patch 内的特征编码,位置顺序由空间网格天然保留,所以直接用 flatten 加 transpose 就能得到 tokens。参数上注意kernel_sizestride必须相等,否则 patch 之间会重叠或出现缝隙,token 数量会和你预期的(image_size // patch) ** 2不一致。

2.2 位置编码加在哪:默认插值方式与维度对齐

Transformer 本身没有空间顺序概念,patch 被拉平成序列后,必须把每个 token 在整个 8x8 网格里的相对位置告诉模型,这就是位置编码的职责。ViT 使用可学习的位置编码,参数维度是[1, num_patches + 1, embed_dim],多出来的 1 是留给 [CLS] token 的。

这里有一个新手必踩的坑:如果从 ImageNet 预训练权重加载,位置编码形状是[1, 197, 384](196 个 patch 加 1 个 [CLS]),而 CIFAR10 用 patch=4 时只需要[1, 65, 384],直接 load 必然报 shape mismatch。常见处置是双线性插值把 196 个位置映射到 64 个位置,但原版位置编码学习的是 224 分辨率下的相对位置语义,插值后的位置顺序并不等价于 32 分辨率下的空间关系,用它初始化有时反而拖慢收敛。我的习惯是从零训练时直接随机初始化,省掉这一层麻烦:

import torch pos_embed = torch.zeros(1, 65, 384) # 64 个 patch + 1 个 [CLS] cls_token = torch.zeros(1, 1, 384) nn.init.trunc_normal_(pos_embed, std=0.02) nn.init.trunc_normal_(cls_token, std=0.02)

trunc_normal_以 0.02 的标准差截断初始化,是 ViT 原论文和大多数开源的默认做法,比常规正态分布更稳,能避免尾部大值在训练初期把 softmax 注意力推向饱和。位置编码在 forward 里和 token 序列直接相加:x = x + pos_embed,相加后维度不变,之后的 Transformer 编码器不需要额外区分哪些维度是位置、哪些是内容。

2.3 分类头与 [CLS] token 的取舍

ViT 原版的做法是在 patch 序列前面拼接一个可学习的 [CLS] token,经过全部 Transformer 层后,取出这个 token 的最终表示送入 MLP 分类头。设计意图是让 [CLS] 通过自注意力聚合整张图的全局信息,类似 BERT 的做法。

但在 CIFAR10 这种小图小数据集上,从零训练时 [CLS] token 不一定是最优选择。实测下来,对所有 patch token 做全局平均池化(GAP)再进分类头,收敛通常更稳,因为梯度可以直接回传到所有 patch,而不是依赖一个 token 在 12 层注意力里逐步聚合信息。数据量越小,这个差异越明显。如果你用的是 224 预训练权重做迁移,建议保留 [CLS] 分支;如果从零训练,我一般把 GAP 作为默认:

class ViT(nn.Module): def __init__(self, ...): ... self.head = nn.Linear(embed_dim, num_classes) def forward_features(self, x): x = self.patch_embed(x) if self.use_cls: cls_token = self.cls_token.expand(x.size(0), -1, -1) x = torch.cat([cls_token, x], dim=1) x = x + self.pos_embed x = self.encoder(x) x = self.norm(x) if self.use_cls: return x[:, 0] # 取 [CLS] token return x.mean(dim=1) # 对所有 patch 做 GAP

加一个use_cls开关,两种模式共用一套代码,切换成本几乎为零。在 CIFAR10 上我会先跑 GAP 版本,如果后续要迁移预训练权重,再把开关打开。

2.4 参数量、计算量与你的显卡预算:选 tiny 还是 small

从零训练 CIFAR10 完全不需要 ViT-Base 这种 8600 万参数的体量,小数据集上参数越多过拟合来得越快。我常用的三档配置如下:

配置depthhiddenheads参数量级单 epoch 参考时间(消费级 GPU)
精简版42564约 2M1 分钟内
默认版63846约 5M1-2 分钟
偏大版65128约 8M2-3 分钟

默认版是大多数情况下最合适的起点:6 层编码器、384 维 hidden、6 个注意力头,参数量在 5M 上下,既能体现 Transformer 的表达能力,又不会让 5 万张训练图撑不住。精简版适合 CPU 上做快速验证,偏大版在小数据集上容易出现训练集准确率接近 100% 而验证集停滞在 80% 左右的过拟合局面。显存紧张时优先砍 hidden 而不是砍 depth,因为 hidden 对参数量和计算量的影响都是二次的,砍 depth 只影响前者。

3. Python 环境与 CIFAR10 数据加载:依赖清单、下载缓存的三个坑、DataLoader 四参数

3.1 依赖装到什么版本够用:torch/torchvision/timm 的兼容组合

先把 python 环境装到位。我建议用 conda 建一个独立环境,python 版本选 3.10 或 3.11,这两个版本对 torch 生态的兼容性最省心。装依赖时注意 torch 和 torchvision 必须配套,pip 直接装官方稳定版会自动解析这对关系,不要手动指定两个不匹配的版本:

pip install torch torchvision pip install matplotlib scikit-learn

torchvision 里带了 CIFAR10 数据集类,是数据加载的核心;matplotlib 用于画损失曲线和混淆矩阵;scikit-learn 在最后做混淆矩阵和 t-SNE 分析时用得上。timm 在这个方案里是可选的:如果你只是想跑通训练和验证,自定义一个几十行的 ViT 就够,不需要它;如果你想对比 timm 里的预训练权重效果,再pip install timm。装完后用一段自检代码确认环境可用:

import torch, torchvision print(torch.__version__, torchvision.__version__) print(torch.cuda.is_available())

cuda.is_available()返回False不代表不能跑,CPU 上也能完成完整的训练和验证,只是每轮会慢 5 到 10 倍。如果没有 GPU,建议把 patch=4 的默认配置换成精简版,并且把后续代码里的 batch size 调小到 64。

3.2 CIFAR10 数据集下载与本地缓存:下载慢、校验失败怎么办

CIFAR10 数据集类的用法很直接,第一次运行会从官方源下载约 170MB 的压缩包并解压到 root 目录:

from torchvision import datasets, transforms transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean=[0.4914, 0.4822, 0.4465], std=[0.2470, 0.2435, 0.2616]) ]) train_dataset = datasets.CIFAR10( root='./data', train=True, transform=transform, download=True) val_dataset = datasets.CIFAR10( root='./data', train=False, transform=transform, download=True)

这段代码里最容易被忽略的是transform里的 Normalize 参数。CIFAR10 的 RGB 三通道均值和标准差就是上面这组数值,不归一化的话,模型输入的数值范围是 0 到 1,而常见预训练或超参调优都是在标准化后的分布上做的,loss 很可能卡在初始值附近不动。

下载环节有三个高频问题。第一,官方源在国内下载很慢甚至超时,常见做法是手动下载cifar-10-python.tar.gz放到./data/目录下,文件名保持原样,然后代码里download=True会检测到文件已存在,跳过下载直接解压。第二,下载中途断开会导致压缩包不完整,下次运行报 checksum mismatch,处置方式是删除./data/下残留的所有文件重新下载,不要只删那个 gz 文件,因为解压出来的缓存目录也可能损坏。第三,如果每次运行都重新下载,检查download参数是否误设成了True且 root 路径不固定,数据集目录在,正常情况第二次运行会直接读缓存。验证数据集只有 1 万张,如果测试集准确率和训练集差太远,先确认训练用的 transform 里有没有加随机增强,验证集是否套用了同样的 Normalize。

3.3 DataLoader 里必须调的 4 个参数:num_workers、pin_memory、drop_last、persistent_workers

DataLoader 的参数看起来简单,调不好会让训练速度差好几倍。我常用的配置如下:

train_loader = torch.utils.data.DataLoader( train_dataset, batch_size=128, shuffle=True, num_workers=4, # Windows 上可以改成 0 或 2,避免子进程卡死 pin_memory=True if torch.cuda.is_available() else False, drop_last=True, persistent_workers=True if num_workers > 0 else False ) val_loader = torch.utils.data.DataLoader( val_dataset, batch_size=128, shuffle=False, num_workers=4, pin_memory=True if torch.cuda.is_available() else False )

四个参数的作用分别是:

参数推荐值作用与坑
num_workersGPU 训练 4-8,CPU 训练 0-2控制数据加载的子进程数。Windows 上开太多容易报多进程初始化错误,建议从 0 起步慢慢加
pin_memory有 GPU 时设 True锁页内存能加速 CPU 到 GPU 的数据拷贝,CPU 训练时设了没意义也没坏处
drop_last训练集设 True,验证集设 FalseCIFAR10 训练集 5 万张除以 128 有余数,不丢弃最后不完整的 batch,训练时 loss 曲线末端会突然跳一段
persistent_workersnum_workers > 0 时设 True每个 epoch 结束后不销毁 worker 进程,省掉重复创建的开销

验证集不要开shuffledrop_last=False,保证 1 万张全部参与统计,准确率才是稳定的。如果你看到每个 epoch 的验证准确率在 1% 到 2% 之间来回跳,先检查是不是没关 shuffle。

4. 从零写 ViT 训练与验证脚本:主循环、验证逻辑与超参

4.1 最小可跑的模型定义与训练主循环

把前面几个组件合并成一个完整的 ViT 类,训练脚本保持在 80 行左右就能跑通:

import torch import torch.nn as nn import math class ViT(nn.Module): def __init__(self, image_size=32, patch=4, num_classes=10, depth=6, num_heads=6, embed_dim=384, mlp_ratio=4): super().__init__() self.patch_embed = PatchEmbed(3, patch, embed_dim) num_patches = (image_size // patch) ** 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)) nn.init.trunc_normal_(self.pos_embed, std=0.02) nn.init.trunc_normal_(self.cls_token, std=0.02) encoder_layer = nn.TransformerEncoderLayer( d_model=embed_dim, nhead=num_heads, dim_feedforward=embed_dim * mlp_ratio, dropout=0.1, activation="gelu", batch_first=True) self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=depth) self.norm = nn.LayerNorm(embed_dim) self.head = nn.Linear(embed_dim, num_classes) def forward_features(self, x): B = x.size(0) x = self.patch_embed(x) # [B, 64, E] cls_token = self.cls_token.expand(B, -1, -1) # [B, 1, E] x = torch.cat([cls_token, x], dim=1) # [B, 65, E] x = x + self.pos_embed x = self.encoder(x) # 6 层 Transformer return self.norm(x[:, 0]) # [CLS] token 特征 def forward(self, x): return self.head(self.forward_features(x)) device = "cuda" if torch.cuda.is_available() else "cpu" model = ViT().to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=5e-2) criterion = nn.CrossEntropyLoss()

nn.TransformerEncoderLayer是 PyTorch 内置的 Transformer 层实现,一个 layer 内部包含多头自注意力、MLP、残差连接和层归一化,batch_first=True让输入输出的形状是[batch, seq, dim],直接和我们的 patch 序列对接,不需要手工 permute。dropout=0.1会同时作用在注意力和 FFN 子层上,防止小数据集过拟合。dim_feedforward=embed_dim * mlp_ratio是 FFN 隐层维度,ViT 惯例放大 4 倍,这里就是 1536。

训练主循环按标准写法组织:

for epoch in range(30): model.train() train_loss, train_correct, train_total = 0.0, 0, 0 for x, y in train_loader: x, y = x.to(device), y.to(device) logits = model(x) loss = criterion(logits, y) optimizer.zero_grad() loss.backward() optimizer.step() train_loss += loss.item() * x.size(0) train_correct += (logits.argmax(dim=1) == y).sum().item() train_total += y.size(0) avg_loss = train_loss / train_total train_acc = train_correct / train_total print(f"epoch {epoch+1:02d} loss {avg_loss:.4f} acc {train_acc:.4f}")

loss.item()取出标量值参与累加,不然会一直保留计算图导致显存累积;logits.argmax(dim=1)拿到每个样本预测的类别下标,和标签比对后累加正确数。optimizer.zero_grad()每次迭代前必须调用,不清零梯度会把多轮的梯度累加到一起,loss 曲线会异常跳动。

4.2 验证循环与准确率统计:不要用训练模式跑验证

验证循环看起来和训练循环很像,但有一个高频错误会导致结果虚高或剧烈抖动:忘记调用model.eval()。ViT 结构里没有 BatchNorm,但 attention 和 FFN 里的 dropout 在训练模式下是激活的,如果你用训练模式跑验证,dropout 会随机丢弃部分神经元,验证准确率每次跑都不一样,忽高忽低。正确写法是:

model.eval() val_loss, val_correct, val_total = 0.0, 0, 0 with torch.no_grad(): for x, y in val_loader: x, y = x.to(device), y.to(device) logits = model(x) loss = criterion(logits, y) val_loss += loss.item() * x.size(0) val_correct += (logits.argmax(dim=1) == y).sum().item() val_total += y.size(0) val_acc = val_correct / val_total print(f"val loss {val_loss / val_total:.4f} acc {val_acc:.4f}")

model.eval()关闭 dropout 相关的随机行为;torch.no_grad()告诉 PyTorch 不要为验证过程构建计算图,既省显存又加速。这两个必须同时存在,只调model.eval()不包no_grad(),显存会被验证过程占掉一大块,训练中途可能 OOM。验证集准确率是你在模型调参过程中的唯一依据,不要看训练集准确率,那里面有 dropout 和模型记忆效应的水分。

4.3 训练超参设置:epoch、学习率、warmup 与 weight decay 怎么给

从零训练 CIFAR10 的 ViT,我常用的超参组合如下:

超参推荐值说明
epochs30-50无数据增强时 30 轮够了,加增强可以往上加
batch_size128CPU 训练减半到 64
lr1e-3AdamW 下这个量级比较稳
weight_decay5e-2和数据增强互补的正则手段
warmup_epochs5学习率从 0 线性升到 1e-3
调度器CosineAnnealingLR训练后期把学习率平滑降到接近 0

Transformer 对学习率很敏感,直接 1e-3 冷启动容易出现前期 loss 震荡甚至发散。我的习惯是前 5 个 epoch 做线性 warmup,之后用余弦退火慢慢降:

def warmup_cosine_lr(step, warmup_steps, total_steps): if step < warmup_steps: return step / max(warmup_steps, 1) progress = (step - warmup_steps) / max(total_steps - warmup_steps, 1) return 0.5 * (1.0 + math.cos(math.pi * progress)) scheduler = torch.optim.lr_scheduler.LambdaLR( optimizer, lr_lambda=lambda step: warmup_cosine_lr(step, 5 * len(train_loader), 30 * len(train_loader)))

LambdaLR的 lambda 接收的是迭代步数而不是 epoch 数,所以要把 epoch 数乘以每个 epoch 的迭代步数换算成总步数。这里的 warmup 步数是 5 个 epoch 对应的步数,total_steps就是 30 个 epoch 的总步数。如果你的 loss 在 warmup 阶段就飘了,把 lr 降到 5e-4 重跑,比调任何结构参数都有效。

5. 训练与验证避坑:5 个让新手翻车的细节

5.1 数据集下载卡住或校验失败

现象:第一次运行卡在下载界面,很久不动;或者中断后再次运行报 checksum mismatch。原因:CIFAR10 官方源在部分网络环境下传输极不稳定,下载的压缩包不完整。解决:手动下载cifar-10-python.tar.gz放到./data/目录下,保持文件名不变,再次运行代码会自动跳过下载;如果已经下载了残缺文件,把整个./data/目录删掉重来,只删 gz 文件不够,解压出来的缓存文件也受影响。这个坑几乎每个人都会遇到,留好一个完整的数据集目录算是最值得做的环境准备。

5.2 loss 卡在 2.3026 附近完全不动

现象:训练了 10 个 epoch,loss 始终在 2.3 左右打转,准确率停留在 10% 附近。原因:2.3026 是 10 分类均匀分布的交叉熵,说明模型完全没学到东西,输出接近等概率。常见诱因有三个:输入没有做 Normalize,数值分布不适合初始化权重;学习率过大导致梯度反复震荡;数据增强过度把图片破坏到无法识别。解决:先确认 transform 里的 Normalize 生效,再检查 lr 是否在 1e-3 量级,最后把数据增强强度降低或临时关掉,等模型能正常收敛再逐步加回来。这个现象在 CNN 上很少见,但 Transformer 结构对输入分布更敏感。

5.3 验证集准确率比训练集还高且波动剧烈

现象:训练集准确率 70%,验证集准确率却显示 80%,而且每轮波动超过 3%。原因:训练模式没有切回验证模式,dropout 在验证时依然生效,相当于每次验证都是随机抽了一部分神经元。解决:验证循环前调用model.eval(),训练循环前调用model.train(),两个方法对应模型里的 dropout 开关。如果你用的是 GPU 训练,还要确认验证时没有把torch.no_grad()写成torch.enable_grad(),那样验证误差会一路累积到爆显存。

5.4 加载预训练权重报 shape mismatch

现象:想从 ViT-Base/16 的 ImageNet 权重迁移到 CIFAR10,load_state_dict报 size mismatch,卡在pos_embedhead上。原因:预训练模型的pos_embed[1, 197, 768],我们的模型是[1, 65, 384],patch 数量和 embedding 维度都不一样。解决:最省心的方案是从零训练,前面的代码不需要加载任何外部权重;如果一定要用预训练权重,把pos_embed做双线性插值后在 main model 之外单独 load,并且分类头head要跳过不加载,因为 ImageNet 是 1000 类,CIFAR10 是 10 类。插值后的位置编码和 32 分辨率下的空间语义不完全对齐,加载后要用小学习率微调,否则前期 loss 会反弹。

5.5 训练到一半内存或显存溢出

现象:训练在第 10 轮左右突然报 CUDA out of memory,或者 CPU 训练时内存占用持续上涨最终被杀进程。原因:最常见的是num_workers开太大,每个 worker 都会复制一份数据集引用;其次是验证循环忘了开no_grad(),把整个验证集的计算图都留在显存里;还有一个隐蔽原因是loss.item()写成了累加 tensor,导致计算图逐步累积。解决:GPU 训练把 batch size 降到 64,验证循环包上torch.no_grad();CPU 训练把num_workers降到 0,数据预读取和主进程竞争内存的情况会明显改善。检查nvidia-smi看显存占用曲线,如果每个 epoch 都在涨,基本就是计算图泄漏,优先排查验证循环。

6. 进阶:用一张混淆矩阵和 t-SNE 看 ViT 在 CIFAR10 上学到了什么

准确率只是一个数字,它不会告诉你模型在哪些类别上混淆。CIFAR10 里猫和狗、鸟和鹿、卡车和汽车是经典难分对,损失一条纹理信息就可能把两类的决策边界搅在一起。我的习惯是每个模型跑完后,第一件事不是对比准确率涨了多少,而是画一张验证集的混淆矩阵:

from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay def collect_preds(model, val_loader, device): model.eval() y_true, y_pred = [], [] with torch.no_grad(): for x, y in val_loader: logits = model(x.to(device)) y_true.extend(y.tolist()) y_pred.extend(logits.argmax(dim=1).cpu().tolist()) return y_true, y_pred y_true, y_pred = collect_preds(model, val_loader, device) cm = confusion_matrix(y_true, y_pred) disp = ConfusionMatrixDisplay(cm, display_labels=val_dataset.classes) disp.plot(cmap="Blues")

混淆矩阵的对角线越亮越好,非对角线上的亮点就是系统性的类别混淆。如果猫和狗之间有一大块亮斑,说明模型学到的是颜色和轮廓层面的一般性特征,而不是能区分的纹理细节,这时候增大 patch 密度或者加强数据增强方向就明确了。

更进一步,可以用 t-SNE 观察模型的倒数第二层特征,也就是forward_features的输出,看 ViT 是否真把 10 个类别在特征空间里分开了。验证集 1 万张全跑会非常慢,常见做法是每个类随机抽 100 张组成 1000 张子集:

from sklearn.manifold import TSNE import matplotlib.pyplot as plt import random random.seed(0) idx = [] for c in range(10): cur = [i for i, t in enumerate(y_true) if t == c] idx.extend(random.sample(cur, 100)) model.eval() feats, labels = [], [] with torch.no_grad(): for i in idx: x, y = val_dataset[i] f = model.forward_features(x.unsqueeze(0).to(device)) feats.append(f.cpu().numpy()[0]) labels.append(y) tsne = TSNE(n_components=2, perplexity=30, random_state=0) xy = tsne.fit_transform(feats) plt.scatter(xy[:, 0], xy[:, 1], c=labels, cmap="tab10", s=8)

t-SNE 的结果需要一些判读经验:理想的图是 10 个点团清晰分离、团内紧凑;如果出现一个大团包着多个类别,通常说明特征还没有充分判别性,优先加训练轮数或加大模型容量;如果点团细碎、同一类别散成好几块,则像是过拟合前期的特征碎片化。判读时留意 t-SNE 的perplexity对密度敏感,30 对 1000 个样本是稳妥起步值,太小会割裂连续聚类。

ViT 在小数据集上很容易表现出一种假象:训练集准确率一路涨到 95%,验证集却卡在 75% 打转。这种局面的后悔药在超参而不是结构,先调数据增强和 weight decay,再考虑换更大模型。我自己的习惯是模型代码写到能用就停,把多出来的时间全花在画图上,毕竟准确率只告诉你答案对不对,混淆矩阵和特征分布才告诉你为什么错。希望帮到你。

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

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

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

立即咨询