拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

ViT在CIFAR10上的训练与验证:Python源码实现与调参实战

ViT在CIFAR10上的训练与验证:Python源码实现与调参实战 简介这份源码包以Vision TransformerViT模型为核心完整实现了CIFAR10图像分类任务的训练与验证流程适合计算机、通信、人工智能、自动化等专业的本科生及从业者作为课程设计、毕业设计或深度学习的入门实践。项目代码经调试可稳定运行包含数据加载、预处理、模型搭建、训练调优和性能评估等模块便于初学者理解Transformer在视觉任务中的应用方式。资源共12个文件以8个Python脚本为主要内容另附README说明、Git配置及训练效果可视化图压缩包仅137KB结构紧凑。目前已有227人学习下载。对希望快速上手ViT分类项目或在此基础上扩展改进的读者来说这套源码提供了清晰的基础框架并保留优化空间具有较高的参考价值。1. 把ViT搬到CIFAR10这套源码解决什么、适合谁用CNN训练CIFAR10分类已经是一件常规到近乎“无脑”的事情但把Vision TransformerViT这种当前ai模型vit主流技术路线搬到32×32的小图上我第一次训练就翻车了loss几乎不下降验证准确率卡在60%上下。回头重写模型、数据增强和学习率策略之后才把准确率稳定在90%以上。“基于Vit实现CIFAR10分类数据集的训练和验证python源码”要解决的正是这整条链路上的问题如何用一份结构清晰的python源码把ViT在CIFAR10上的训练与验证流程完整跑通而不是停留在调通一个demo。这篇文章按“原理→代码→参数→踩坑”展开适合从CNN转向Transformer的开发者、准备课程设计或论文对比实验的同学也适合想判断vit结构在小图分类任务上是否值得投入的算法工程师。2. ViT结构拆解与CIFAR10适配为什么patch_size4是第一个关键决定2.1 ViT核心结构patch embedding、transformer encoder、分类头标准ViT的第一步是把图像切块patch。对32×32的CIFAR10图像如果沿用ImageNet上的patch_size16一张图只能切成4个patch序列长度太短注意力机制几乎没有空间信息可用视觉Transformer几乎退化成全连接。很多初学者直接套vit_base_patch16_224结果验证集准确率始终上不去。正确的做法是把patch_size改为4让每个patch覆盖4×4像素图像被切成8×8共64个patch序列长度64再加一个分类token完整序列长度65。这个数值几乎决定了后续所有shape判断。patch embedding在主流源码里用卷积实现class PatchEmbed(nn.Module): 把图像切成patch并线性投影到embedding空间 参数说明 img_size: 输入图像边长, CIFAR10固定为32 patch_size: patch边长, 小图任务取4或8, 取16会退化成4个patch in_chans: 输入通道数, RGB是3 embed_dim: 每个patch投影后的特征维度, 常见取192/256/384 def __init__(self, img_size32, patch_size4, in_chans3, embed_dim192): super().__init__() self.num_patches (img_size // patch_size) ** 2 # 32//48, 8*864 self.proj nn.Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): # x形状: [B, 3, 32, 32] x self.proj(x) # [B, 192, 8, 8] x x.flatten(2) # [B, 192, 64] x x.transpose(1, 2) # [B, 64, 192] return x这段代码用kernel_sizepatch_size且stridepatch_size的卷积完成“切块线性投影”每个4×4×3的patch被压缩成192维向量。用Conv2d实现的好处是走矩阵乘法和GPU并行比逐patch取数再进Linear快很多这也是timm等主流库的通用写法。flatten和transpose共同把通道维挪到最后得到Transformer需要的[B, seq_len, embed_dim]形状。切块之后在序列最前面拼接一个class token。这个token是可学习的向量最终分类只用它对应的输出而不对64个patch的输出做池化cls_tokens self.cls_token.expand(B, -1, -1) # [B, 1, embed_dim] x torch.cat([cls_tokens, x], dim1) # [B, 65, embed_dim]位置编码在这之后直接加到token序列上。self-attention本身不感知顺序位置编码是模型知道“patch在第几行第几列”的主要途径self.pos_embed nn.Parameter(torch.zeros(1, 65, embed_dim)) x x self.pos_embedCIFAR10输入尺寸固定可学习位置编码比sinusoid更省心不需要任何插值逻辑。再往后就是重复堆叠的TransformerEncoder block先LayerNorm再做多头自注意力再做LayerNorm和MLP每个子层带残差连接分类头就是一层Linear把class token对应的192维向量映射到10类。理解了这个骨架后面改参数时就清楚每个旋钮影响的是哪一段。2.2 为什么经典ViT在CIFAR10上效果不理想patch、数据量和正则如果说patch_size4是源码里的第一个关键决定那第二个关键决定是不要直接加载ImageNet预训练权重。最常见的问题是shape对不上224分辨率、patch16的预训练模型序列长度为196位置编码是[1, 197, 768]而CIFAR10下patch4的序列长度只有64位置编码是[1, 65, 192]embed_dim也不同。就算手动插值位置编码预训练学到的分辨率相关特征迁移到32×32小图上往往是负资产。CIFAR10对这种规模的任务来说从零训练配合合理的增强与正则完全够用。抛开权重迁移问题ViT本身在小数据集上有一个天然劣势它不像CNN自带局部性和平移等变性。CNN在以很少的数据学习边缘纹理时卷积核天然只在局部窗口内滑动而self-attention一上来就在全部65个token之间做注意力要从头学“哪些token关系有用”。CIFAR10训练集只有5万张图对ViT这类缺少归纳偏置的结构偏少所以源码里通常要补两种手段增强数据增强RandomCrop、RandomHorizontalFlip只是基础CutMix或MixUp能显著提升ViT最终精度加强正则weight_decay设到0.05或0.1配合dropout和warmup。如果这些不做直接照搬CNN超参训练100轮往往只有75%到85%而同预算的ResNet已经能到93%以上。这也是我不建议拿ViT和ResNet在CIFAR10上做简单“暴力对比”的原因——增强和正则没有对齐结论没有参考价值。2.3 cifar10数据集下载与DataLoader的transform参数配置CIFAR10数据集下载最常见的方式是torchvision自带接口第一次运行会自动下载并解压到指定目录之后走缓存不再重复下载。注意Mean和Std是官方统计的标准值不要自己重新统计否则归一化结果和预训练约定不一致import torch from torchvision import datasets, transforms # CIFAR10训练集和测试集通用的归一化参数 CIFAR10_MEAN (0.4914, 0.4822, 0.4465) CIFAR10_STD (0.2470, 0.2435, 0.2616) transform_train transforms.Compose([ transforms.RandomCrop(32, padding4), # 四周补4像素再随机裁剪, 缓解边缘敏感 transforms.RandomHorizontalFlip(), # 随机水平翻转, 对CIFAR10无方向性任务有效 transforms.ToTensor(), transforms.Normalize(CIFAR10_MEAN, CIFAR10_STD), ]) transform_test transforms.Compose([ transforms.ToTensor(), # 验证集不做增广, 只归一化 transforms.Normalize(CIFAR10_MEAN, CIFAR10_STD), ]) trainset datasets.CIFAR10( root./data, trainTrue, downloadTrue, transformtransform_train ) testset datasets.CIFAR10( root./data, trainFalse, downloadTrue, transformtransform_test ) trainloader torch.utils.data.DataLoader( trainset, batch_size128, shuffleTrue, num_workers4, pin_memoryTrue, ) testloader torch.utils.data.DataLoader( testset, batch_size256, shuffleFalse, num_workers4, pin_memoryTrue, )这里有几个新手容易忽略的参数。num_workers4让数据读取和模型计算并行不要设成0否则GPU会在每个batch前空等。pin_memoryTrue配合后面的images.cuda(non_blockingTrue)能减少CPU到GPU的拷贝阻塞。testloader的shuffleFalse很重要——验证时不需要打乱顺序而且只有保持固定顺序后面做混淆矩阵时才能把模型输出和真实标签严格对齐。下载完成后建议先跑一个shape检查取一个batch确认模型中间输出是[B, 65, 192]最终输出是[B, 10]这一步能拦截大部分维度不对齐的问题。3. 从零搭训练和验证代码模型定义、训练循环与checkpoint保存3.1 手写ViT还是timm两条路线怎么选先回答最常被问的问题要不要用timm。我的建议是验证算法思路用timm改结构、做课程设计或想深入读源码就手写。timm的vit_tiny_patch16_224、vit_base_patch16_224调用简单但默认ImageNet分辨率想在CIFAR10上跑必须改input尺寸、patch数量和位置编码绕不开源码修改。手写一份不依赖timm的紧凑实现环境依赖更少跑起来也更接近最小可用的训练验证框架。下面是完整模型定义核心三部分PatchEmbed、TransformerBlock、ViT主体。这里把embed_dim设为192depth为6num_heads为6在CIFAR10上是性价比比较高的配置参数量约200万上下比vit_base小一个量级单卡就能跑完100轮。class TransformerBlock(nn.Module): def __init__(self, dim, num_heads, mlp_ratio4.0, dropout0.1): super().__init__() self.norm1 nn.LayerNorm(dim) self.attn nn.MultiheadAttention( dim, num_heads, dropoutdropout, batch_firstTrue ) self.norm2 nn.LayerNorm(dim) self.mlp nn.Sequential( nn.Linear(dim, int(dim * mlp_ratio)), nn.GELU(), nn.Dropout(dropout), nn.Linear(int(dim * mlp_ratio), dim), nn.Dropout(dropout), ) def forward(self, x): # PreNorm写法: 先归一化再进注意力, 再残差连接 hn self.norm1(x) x x self.attn(hn, hn, hn, need_weightsFalse)[0] x x self.mlp(self.norm2(x)) # 残差连接 return x关键点有两个。第一nn.MultiheadAttention用batch_firstTrue输入形状是[B, seq_len, embed_dim]返回值是一个tuple取第0个元素才是注意力输出。need_weightsFalse可以跳过注意力权重矩阵的计算省显存也省时间。第二按上面的写法norm1只调用了一次先用hn保存归一化结果再喂给注意力避免重复计算这是对论文结构图的一种效率优化输出不变。接着是ViT主体class ViTForCIFAR10(nn.Module): def __init__(self, img_size32, patch_size4, in_chans3, num_classes10, embed_dim192, depth6, num_heads6, mlp_ratio4.0, dropout0.1): super().__init__() self.patch_embed PatchEmbed(img_size, patch_size, in_chans, embed_dim) self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed nn.Parameter( torch.zeros(1, self.patch_embed.num_patches 1, embed_dim) ) self.pos_drop nn.Dropout(pdropout) self.blocks nn.Sequential(*[ TransformerBlock(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) self._init_weights() def _init_weights(self): # 位置编码和分类token都用截断正态初始化, 初始方差别太大 nn.init.trunc_normal_(self.pos_embed, std0.02) nn.init.trunc_normal_(self.cls_token, std0.02) for m in self.modules(): if isinstance(m, nn.Linear): nn.init.trunc_normal_(m.weight, std0.02) if m.bias is not None: nn.init.zeros_(m.bias) def forward(self, x): B x.shape[0] x self.patch_embed(x) # [B, 64, 192] cls_tokens self.cls_token.expand(B, -1, -1) x torch.cat([cls_tokens, x], dim1) # [B, 65, 192] x x self.pos_embed x self.pos_drop(x) x self.blocks(x) # [B, 65, 192] x self.norm(x) return self.head(x[:, 0]) # 只取class token分类初始化往往被新手跳过但ViT对初始化比CNN敏感。位置编码初始值太大前几个epoch注意力会偏向固定位置收敛变慢trunc_normal_配合std0.02是timm和官方ViT源码的标准做法。embed_dim192配num_heads6每个注意力头分到32维符合d_k通常在32到64之间的经验区间如果头数多而通道少每个头的表达能力会不够。3.2 训练和验证循环先写对再优化训练循环和验证循环建议写成两个独立函数。训练循环里必须model.train()验证循环里必须model.eval()这个切换决定dropout和LayerNorm的行为。def train_one_epoch(model, trainloader, criterion, optimizer, epoch): model.train() run_loss, run_correct, run_total 0.0, 0, 0 for images, labels in trainloader: images images.cuda(non_blockingTrue) labels labels.cuda(non_blockingTrue) optimizer.zero_grad() outputs model(images) # [B, 10] loss criterion(outputs, labels) loss.backward() optimizer.step() run_loss loss.item() * images.size(0) _, predicted outputs.max(dim1) # 取最大logit对应类别 run_total labels.size(0) run_correct predicted.eq(labels).sum().item() return run_loss / run_total, run_correct / run_total这段代码里run_loss按样本数加权累加最后除以样本总数避免最后一个batch不满时统计偏斜。outputs.max(dim1)返回最大值和对应下标用_丢弃最大值用predicted拿类别索引。验证循环结构几乎相同但不需要backward和optimizer并且建议加torch.no_grad()装饰器避免无谓构建计算图:torch.no_grad() def validate(model, testloader, criterion): model.eval() run_loss, run_correct, run_total 0.0, 0, 0 for images, labels in testloader: images images.cuda(non_blockingTrue) labels labels.cuda(non_blockingTrue) outputs model(images) loss criterion(outputs, labels) run_loss loss.item() * images.size(0) _, predicted outputs.max(dim1) run_total labels.size(0) run_correct predicted.eq(labels).sum().item() return run_loss / run_total, run_correct / run_total如果你看到报错expected scalar type Long but found Float多半是labels类型问题比如忘了把标签从long转成float的模型输入。如果你看到维度对不上先检查是不是忘了取x[:, 0]就把整个序列丢给了分类头。3.3 checkpoint保存与断点续训别让一次断电重跑40轮ViT训练动辄几十上百轮不保存checkpoint就是拿一天时间赌运气。源码里的策略一般是按验证准确率保存最优模型而不是保存最后一个epoch。Transformer训练中后期验证集准确率曲线会震荡最后一个epoch不一定是泛化最好的。best_acc 0.0 for epoch in range(1, total_epochs 1): train_loss, train_acc train_one_epoch( model, trainloader, criterion, optimizer, epoch ) val_loss, val_acc validate(model, testloader, criterion) if scheduler is not None: scheduler.step() print(fEpoch {epoch:3d} | Train {train_acc:.4f} | fVal {val_acc:.4f} | lr {optimizer.param_groups[0][lr]:.2e}) if val_acc best_acc: best_acc val_acc torch.save({ model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), epoch: epoch, best_acc: best_acc, scheduler_state_dict: scheduler.state_dict() if scheduler else None, }, vit_cifar10_best.pth) print(fsave checkpoint, best acc {best_acc:.4f})断点续训是这里最容易被忽略的一环。由于只保存了model和optimizer的state_dict恢复时需要先重建模型再注入权重并把epoch和scheduler状态一并恢复。cosine调度器如果从头算学习率曲线会错乱def resume_training(model, optimizer, checkpoint_path, schedulerNone): ckpt torch.load(checkpoint_path, map_locationcuda) model.load_state_dict(ckpt[model_state_dict]) optimizer.load_state_dict(ckpt[optimizer_state_dict]) start_epoch ckpt[epoch] 1 if scheduler is not None and ckpt.get(scheduler_state_dict) is not None: scheduler.load_state_dict(ckpt[scheduler_state_dict]) return model, optimizer, scheduler, start_epoch, ckpt[best_acc]保存scheduler是很多源码容易漏掉的点。如果恢复时只恢复optimizer而不恢复scheduler学习率会被重置到起点附近续训后的第一个epoch loss会突然跳高然后慢慢回落白白浪费几个epoch。4. 调参实战让ViT在CIFAR10上正常收敛的必改参数4.1 学习率与warmupAdamW、1e-3、少量epoch预热ViT的学习率经验值和CNN不太一样。CIFAR10上我用得最稳的配置是AdamWlr1e-3weight_decay0.05。按CNN习惯设成1e-4会出问题前十几个epoch loss几乎不动因为ViT的分类头和位置编码是随机初始化的1e-4这个步长对Transformer来说太小。反过来设成5e-3又容易在前几轮发散loss直接变NaN。1e-3配合warmup是一个比较平衡的起点。warmup用LambdaLR实现import math from torch.optim.lr_scheduler import LambdaLR def build_scheduler(optimizer, warmup_epochs5, total_epochs100): def lr_lambda(epoch): if epoch warmup_epochs: return (epoch 1) / warmup_epochs progress (epoch - warmup_epochs) / max(1, total_epochs - warmup_epochs) return 0.5 * (1.0 math.cos(math.pi * progress)) return LambdaLR(optimizer, lr_lambda)注意LambdaLR传入的epoch从0开始第一个epoch时lr是base_lr的1/5不是base_lr。有人把epoch从1开始计数warmup的第一轮反而拿到更高lr效果会受影响。如果你用了断点续训务必把scheduler也保存和恢复否则warmup计次清零lr会重新爬坡。还有一个容易被忽略的细节position embedding和cls token究竟该不该做weight decay。timm会按参数名过滤手写实现时最省事的做法是对所有参数做decay但weight_decay不要超过0.1。如果追求精细可以只对Linear和Conv2d参数做decay不对位置编码做。这两种做法在CIFAR10上的差距不算大我更倾向于先跑通再优化。4.2 训练轮数、batch size与数据增强的配合100轮与早停CIFAR10上有个三角关系batch size越大训练越稳当但大批量下小数据集的泛化损失越明显数据增强越强达到目标精度所需的epoch越多但最终精度更高epoch越多后段过拟合风险越高。下面这组参数是多次实践里比较稳的起点参数建议值说明batch size128显存2G以上能跑4G以上推荐256embed_dim192提升到256收益有限显存增加明显depth6加深到8级在CIFAR10上收益不大epochs100少于60轮很难摸到90%warmup epochs5占训练周期的5%左右weight_decay0.050.05~0.1区间效果不错dropout0.1小模型取0.1即可过大会欠拟合增强RandomCropFlip追求更高精度可加MixUp或CutMix有个常见误区是把epoch设成300想当然以为更久更好。CIFAR10上跑到中后段验证准确率曲线基本走平训练loss还在缓慢下降多出来的是过拟合不是泛化收益。我在跑这个任务时用100轮配合早停早停条件不是“准确率不涨”而是“验证loss连续15轮不降”。只用准确率判断容易误判因为90%以上时准确率变化很钝几个epoch可能都停在90.2%附近但验证loss可能已经从0.28涨到0.34提前预警过拟合。4.3 训练与验证的日志里该看哪些信号训练中值得同时盯四个信号train loss、train acc、val loss、val acc。train loss前5轮如果下降明显慢于预期先怀疑lr和warmuptrain loss和val loss的gap超过0.8基本说明过拟合val acc到达90%以后val loss比val acc更敏感更适合做早停和模型选择的依据。我习惯每轮结束都打印一行完整日志格式固定print(fEpoch {epoch:3d} | fTrain Loss {train_loss:.4f} Acc {train_acc:.4f} | fVal Loss {val_loss:.4f} Acc {val_acc:.4f} | flr {optimizer.param_groups[0][lr]:.2e})这行日志的价值在于它能串起学习率、训练集表现和验证集表现三者的关系。比如某一轮val acc掉点你可以立刻回看lr走到了cosine曲线的哪个位置如果正好是lr快速下降的中段那掉点可能只是正常的调度震荡就不需要急着调参。很多翻车现场都是因为不看日志、只盯最终准确率最后连什么时候开始过拟合的都说不清楚。5. ViT训练CIFAR10避坑清单5个能省半天的常见问题5.1 训练loss卡在2.3附近不动现象loss在2.3附近震荡怎么训练都不降验证准确率稳定在10%左右接近均匀随机猜测。原因2.3正是CIFAR10十类均匀分布的交叉熵−log(0.1)≈2.3026。模型完全没有学到有效特征。最常见的学习率过低或者warmup阶段lr被压得太低另一种可能是初始化时位置编码方差过大导致前几个epoch注意力分布极端。解决先确认optimizer的param_groups里lr真的是1e-3不是被warmup降成了1e-4级别打印lr日志第一个epoch的lr应该是base_lr/warmup_epochs不由base_lr直接决定。如果lr没问题检查初始化把pos_embed和cls_token重新用std0.02的trunc_normal初始化再训练。初始化问题在小数据集上特别容易被误判为超参问题。5.2 加载预训练位置编码时shape mismatch现象想加载vit_base_patch16_224的权重做迁移学习torch.load之后model.load_state_dict直接报错size mismatch for pos_embedtorch.Size([1, 197, 768]) vs torch.Size([1, 65, 192])。原因预训练模型序列长度是196个patch加1个cls token共197embed_dim是768而CIFAR10上的ViT序列长度为65embed_dim只有192。形状完全不同直接加载必然报错。解决两条路。一是放弃预训练数据增强和warmup到位后从零训练到90%左右并不难二是坚持迁移就要对位置编码做二维插值并只能加载能对齐的层。但插值比例过大时预训练学到的相对位置语义会被破坏反而损失精度。CIFAR10规模的任务我建议直接走从零训练省掉各种花式麻烦。5.3 验证集准确率波动大最后一步还突然掉点现象val acc曲线在85%到90%之间来回摆动保存的best模型表现不错但最后一轮之后的val acc反而不如best。原因训练后期lr已经退火到很低val acc在小范围内震荡属于正常现象。但如果每轮都保存last状态不小本文还有配套的精品资源点击获取
返回列表