ARTICLE DETAIL

资讯详情

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

ViT源码逐行解析:从Patch Embedding到人脸情感识别

ViT源码逐行解析:从Patch Embedding到人脸情感识别 第一次把 ViTVision Transformer的源码拖到本地跑通我盯着控制台里那行(1, 197, 768)的输出愣了半天——一张 224×224 的 RGB 人脸图怎么就变成了一串 197 个、每个 768 维的向量cls token 是从哪冒出来的位置编码凭什么能直接相加而不是拼接如果你也卡在这个阶段这篇 Vision Transformer 代码分析就是写给你的。我不打算复述论文里的公式而是把 VIT 代码从输入张量到最终情感标签的整条链路拆开逐行讲清楚每个操作的形状变化、设计动机和实际调试时的坑。内容覆盖 Patch Embedding、多头自注意力、Pre-LN 结构、预训练权重加载、显存优化以及把它迁移到人脸图像情感识别任务时的改造清单同时对比一下 ViT 和 EfficientNetV2 在中小数据集上的选型取舍。适合已经看得懂 PyTorch 基础 API、想真正把 ViT 跑起来并改造成自己项目的人。1. 用数据流视角重看ViT一张人脸图是怎么变成标签的大多数人读 ViT 源码失败不是因为某个算子看不懂而是脑子里没有一条完整的张量流动路径。论文图看十遍不如自己打印一次中间形状。这一节我按数据实际流经的顺序把整条链路串一遍后面几节再逐个模块展开。1.1 输入阶段的三个约定先定死再谈模型ViT 对输入有三个隐含约定代码里通常写在配置类里但新手最容易忽略图像必须被 resize 成固定尺寸比如 224×224。原因在于位置编码的参数量是num_patches 1它和输入分辨率是强绑定的输入变了位置编码就对不上。通道数默认是 3。做人脸情感识别时如果你手上的数据集是 FER2013 那种 48×48 灰度图就得先复制成三通道或者把in_chans改成 1 并重新初始化第一层卷积。Batch 维度必须在最前面即(B, C, H, W)。ViT 的实现里很少做形状兼容判断喂进去(C, H, W)会直接报维度错误。我给初学者一个自检习惯在forward的第一行加上print(x.shape)跑一次单样本。这一步花十秒能省掉后面半小时的排错。1.2 Patch Embedding这一步到底在做什么Patch Embedding 是 ViT 最反直觉的地方。它把一个图像切块、拉平、线性映射三步操作用一个步长等于卷积核大小的二维卷积就完成了。以patch_size16、img_size224、embed_dim768为例卷积核 16×16、步长 16输出特征图尺寸是 14×14也就是 196 个位置每个位置的通道数是 768恰好等于embed_dim展平后得到 196 个 token每个 768 维。为什么用卷积而不是unfold或者切片因为卷积在底层调用的是高度优化的矩阵乘实现速度比手工切片快得多而且参数量完全一致16×16×3×768。这不是为了优雅是实打实的性能选择。提示卷积核大小必须等于步长否则 patch 会重叠或漏掉像素。如果你把patch_size设成 14 而img_size是 224224/1416能整除没问题但如果设成 12会出现除不尽的情况位置编码数量就要手动算建议直接用能整除的数值。1.3 cls token、位置编码和 Encoder 的串联顺序张量进入 Encoder 之前会做三件事顺序不能乱拼接 cls token一个可学习的向量形状(1, 1, 768)用expand复制到整个 batch。它不来自任何图像区域作用相当于一个全局信息汇聚槽。加位置编码注意是加法不是拼接。这相当于给每个 token 打上一个坐标水印因为自注意力本身对顺序不敏感没有位置信息的话打乱 patch 顺序结果完全一样。接 Dropout训练时随机丢弃一部分 token 的激活值防止过拟合。拼完 cls token 后token 数量从 196 变成 197这也是为什么你看到(1, 197, 768)这个形状。分类时只取第 0 个位置cls token的输出喂给分类头其余 196 个 token 的特征被丢弃。这个设计比平均池化的效果在多数任务上更稳定因为它让网络自己学会往哪个位置汇总信息。2. PatchEmbed与多头注意力的逐行实现拆解第一节讲的是流这一节讲零件。我把两个最容易写错、也最值得深挖的模块拿出来一行一行看。理解这两个剩下的 MLP、LayerNorm 基本就是填空题。2.1 卷积切块的等价性用代码验证一遍很多人对卷积等于切块全连接这件事只停留在口头理解。我给一段可以跑的最小验证代码跑完之后你对 Patch Embedding 就不会再有疑惑import torch import torch.nn as nn torch.manual_seed(0) x torch.randn(1, 3, 32, 32) patch_size, embed_dim 8, 64 # 方式一卷积实现 conv nn.Conv2d(3, embed_dim, kernel_sizepatch_size, stridepatch_size) out_conv conv(x).flatten(2).transpose(1, 2) # (1, 16, 64) print(conv:, out_conv.shape) # 方式二手工切块 线性层 unfold nn.Unfold(kernel_sizepatch_size, stridepatch_size) patches unfold(x) # (1, 3*8*8, 16) linear nn.Linear(3 * patch_size * patch_size, embed_dim) out_manual linear(patches.transpose(1, 2)) # (1, 16, 64) print(manual:, out_manual.shape) # 把卷积权重重排后赋给线性层两者输出应当完全一致 linear.weight.data conv.weight.data.view(embed_dim, -1) linear.bias.data conv.bias.data.clone() print(max diff:, (out_conv - out_manual).abs().max().item())实测下来max diff会是一个 1e-6 量级的值纯粹是浮点误差。看到这个数字你就知道两种写法在数学上是一回事区别只在工程效率。2.2 QKV 投影与 reshape/transpose 的维度陷阱多头自注意力最让人头大的是那句先 reshape 再 permute。我把它拆成五步每步标注形状以B2, N197, C768, heads12为例B, N, C 2, 197, 768 heads 12 head_dim C // heads # 64 qkv self.qkv(x) # (B, N, 3*C) - (2, 197, 2304) qkv qkv.reshape(B, N, 3, heads, head_dim) # (2, 197, 3, 12, 64) qkv qkv.permute(2, 0, 3, 1, 4) # (3, 2, 12, 197, 64) q, k, v qkv[0], qkv[1], qkv[2] # 各 (2, 12, 197, 64) attn (q k.transpose(-2, -1)) * self.scale # (2, 12, 197, 197) attn attn.softmax(dim-1) out (attn v) # (2, 12, 197, 64) out out.transpose(1, 2).reshape(B, N, C) # (2, 197, 768)这里有两个必须记住的点。第一permute 的顺序是 (2, 0, 3, 1, 4)目的是把3这个 QKV 维提到最前面把 heads 维提到 batch 后面让后面的矩阵乘默认在最后两维上做。第二q k.transpose(-2, -1)得到的是(N, N)的注意力矩阵197×197 里每一行代表某个 token 对包括自己在内所有 token 的关注程度。self.scale为什么是head_dim ** -0.5因为 Q 和 K 的点积随维度增大而方差变大数值会膨胀softmax 之后会变得极其尖锐梯度趋近于零。除以维度的平方根本质上是把点积的方差拉回到 1 附近让梯度保持健康。2.3 Pre-LN、残差与 DropPath 的三种写法差异Transformer Block 的结构顺序不同实现差异很大这也是读别人代码时最容易困惑的地方结构操作顺序特点常见出处Post-LNx x Attn(LN(x))中的 LN 放在后面原始 Transformer 用法训练深层网络需要精细调 warmup早期实现Pre-LNx x Attn(LN(x))中的 LN 放在分支内前部梯度更稳深层也能直接训ViT 主流实现残差缩放分支输出乘以一个小于 1 的常数抑制深层激活值爆炸大模型常用ViT 官方实现用的是 Pre-LN也就是x x drop_path(attn(norm1(x)))接着x x drop_path(mlp(norm2(x)))。这样写的好处是残差路径上没有任何归一化操作梯度可以沿着恒等映射一路回传。DropPath也叫随机深度是另一个容易被忽略的细节它不是 dropout 激活值而是在训练时随机把整个 Block 的输出置零并按比例缩放。对于 ViT-Base 这种 12 层的结构drop path rate 通常从 0 线性增加到 0.1浅层几乎不丢深层丢得多。这个策略在小数据集微调时特别管用。class DropPath(nn.Module): def __init__(self, drop_prob0.0): super().__init__() self.drop_prob drop_prob def forward(self, x): if self.drop_prob 0.0 or not self.training: return x keep 1 - self.drop_prob shape (x.shape[0],) (1,) * (x.ndim - 1) mask x.new_empty(shape).bernoulli_(keep) return x * mask / keep注意最后那句x * mask / keep除以保留概率是为了保证期望值不变。少了这一步推理时的数值尺度会和训练时对不上表现就是验证集精度莫名其妙掉几个点。3. 迁移到人脸图像情感识别时的改造清单代码能跑通只是起点。真正落地到基于深度学习的人脸图像情感识别这类任务时ViT 需要做几处针对性改造。这一节我把改造点和背后的理由讲清楚你照着改基本不会走偏。3.1 分辨率、通道与分类头的三处必然修改先看几个参数怎么定。假设你的数据集是常见的 7 类情感愤怒、厌恶、恐惧、开心、悲伤、惊讶、中性输入分辨率如果原图是 48×48 的灰度图直接上patch_size16只能切出 3×39 个 patch信息量太少。我的做法是双线性插值放大到 224×224patch_size 保持 16得到 196 个 token。通道数灰度图要么复制三份伪造成 RGB要么把 PatchEmbed 的in_chans改成 1。前者可以直接复用预训练权重后者必须重新初始化第一层卷积会丢掉预训练的纹理先验。小数据集上我建议前者。分类头把原来的nn.Linear(768, 1000)换成nn.Linear(768, 7)。如果数据量小于一万张最好在分类头前再插一个Dropout(0.3)或者干脆只训练最后几层。class ViTForEmotion(nn.Module): def __init__(self, backbone, num_classes7, dropout0.3): super().__init__() self.backbone backbone self.head nn.Sequential( nn.LayerNorm(768), nn.Dropout(dropout), nn.Linear(768, num_classes), ) def forward(self, x): # backbone 返回 (B, N1, C)取 cls token feats self.backbone(x)[:, 0] return self.head(feats)注意换分类头之后新加的层是随机初始化的如果和预训练权重用同一个学习率很容易把已经学好的特征带偏。常见做法是把 backbone 的学习率设成分类头的 1/10。3.2 ViT 与 EfficientNetV2 的取舍别只看论文指标热词里同时出现了 ViT 和 EfficientNetV2这两个确实是当前图像分类绕不开的选项但它们的适用场景差别很大。我按实际项目经验整理了一张对照表维度Vision TransformerEfficientNetV2归纳偏置几乎没有完全靠数据学强卷积自带局部性和平移不变性数据需求大百万级起步才发挥优势小几千到几万张就能训得像样参数效率ViT-Base 约 86MEfficientNetV2-S 约 21M推理速度中等注意力是 O(N²)更快尤其在小分辨率下小数据集表现容易过拟合必须靠强增强和预训练相对稳开箱即用可解释性注意力图可视化直观需要 Grad-CAM 之类的工具我的判断标准很简单数据量低于五万张先跑 EfficientNetV2 做 baseline数据量足够或者有大规模预训练权重可用再上 ViT。人脸情感识别这类数据集普遍偏小所以我一般会同时跑两个用 ViT 做上限探索用 EfficientNetV2 做稳妥交付。3.3 小样本过拟合的对抗手段从数据侧先动手ViT 在几千张图上训练通常三个 epoch 就开始过拟合。我的处理顺序是先从数据侧下手再动模型强增强组合RandAugment Random Erasing Mixup/CutMix。Mixup 对 ViT 尤其有效因为它能让注意力分布更平滑。标签平滑label_smoothing0.1情感识别里存在标注歧义比如惊讶和恐惧表情很接近标签平滑能缓解这个问题。分层冻结先冻结前 8 层只训后 4 层和分类头跑 5 个 epoch 后再全部解冻用更小的学习率微调。类别重采样情感数据集通常开心样本远多于厌恶用WeightedRandomSampler平衡一下比调模型结构有效得多。还有一个经验别急着上大模型。ViT-Tinyembed_dim1926 层在人脸情感这种细粒度分类上往往比 ViT-Base 表现更好因为它不容易在少量数据上记住噪声。4. 我踩过的那些坑从形状报错到训练不收敛的排查链路这一节是整篇最花时间写的部分。下面这些坑我都真实踩过排错过程也按当时的思路还原你可以直接对照自己的报错信息定位。4.1 形状不匹配报错先看是哪一层炸的最常见的一类报错长这样RuntimeError: mat1 and mat2 shapes cannot be multiplied (28x768 and 1024x768)这句话的关键信息是28x768说明参与矩阵乘的张量第一维是 28。28 是 7×4 还是别的如果 batch_size 是 1那很可能是你的 patch 数量算错了。排查顺序打印 PatchEmbed 的输出形状确认num_patches是不是(H//patch_size) * (W//patch_size)检查位置编码的num_patches是否和实际一致如果位置编码是nn.Parameter(torch.zeros(1, 197, 768))而实际 token 数是 197那就没问题如果是 145对应 12×12就会在加法那一步报形状错误。另一类高频报错是加载预训练权重时size mismatch for pos_embed: copying a param with shape [1,197,768] from checkpoint, the shape in current model is [1,50,768].这说明你的输入分辨率或者 patch_size 和预训练时不一致。解决办法是对位置编码做双线性插值公式是把原始的 14×14 网格重排成二维插值到目标网格大小再展平回去def interpolate_pos_embed(pos_embed, old_grid14, new_grid7): cls_token pos_embed[:, :1] # (1, 1, C) patches pos_embed[:, 1:] # (1, 196, C) dim patches.shape[-1] patches patches.reshape(1, old_grid, old_grid, dim).permute(0, 3, 1, 2) patches torch.nn.functional.interpolate( patches, size(new_grid, new_grid), modebicubic, align_cornersFalse ) patches patches.permute(0, 2, 3, 1).reshape(1, new_grid * new_grid, dim) return torch.cat([cls_token, patches], dim1)注意这里插值的对象是特征维度为通道的二维图所以要先 permute 再 interpolate 再 permute 回来顺序错了形状会莫名其妙。4.2 权重加载时 cls_token 和 pos_embed 的特殊处理预训练权重里cls_token和pos_embed是两个独立参数不在任何子模块里。用load_state_dict(strictFalse)时它们会被静默忽略模型照跑但精度上不去——这是个很隐蔽的坑因为不报错。我的做法是写一个显式的加载函数把关键参数逐个核对def load_pretrained(model, ckpt_path, verboseTrue): ckpt torch.load(ckpt_path, map_locationcpu) state ckpt.get(model, ckpt) msg model.load_state_dict(state, strictFalse) if verbose: print(missing:, msg.missing_keys) print(unexpected:, msg.unexpected_keys) assert cls_token not in msg.missing_keys, cls_token 没加载上检查键名 assert pos_embed not in msg.missing_keys, pos_embed 没加载上检查键名 return model用断言把这两个键卡死比事后靠精度异常去反查效率高得多。另外如果换了分辨率记得先插值再加载。4.3 显存爆炸与 loss 不下降的两种典型场景显存问题在 ViT 上比 CNN 突出得多因为注意力矩阵是 O(N²)。224 分辨率下 197×197 看起来不大但如果 batch_size 开到 64、12 个 head中间激活会迅速吃满显存。我的排查顺序先用batch_size1跑通前向和反向确认模型本身没泄漏打开torch.cuda.amp.autocast()混合精度通常能省 30%~40% 显存用梯度累积代替大 batchaccum_steps4, batch_size16等效于 64最后才是考虑减小分辨率或者换 ViT-Tiny。至于 loss 不下降我遇到过三次原因各不相同第一次是学习率设成了 1e-3对于微调来说太高改成 1e-5 立刻就动了第二次是忘了加 warmup前几百步梯度爆炸第三次最隐蔽——位置编码被初始化成零而不是论文要求的截断正态分布导致模型完全学不到空间信息。这三条我建议你都检查一遍。# 位置编码的正确初始化方式 nn.init.trunc_normal_(self.pos_embed, std0.02) nn.init.trunc_normal_(self.cls_token, std0.02)5. 训练流程的工程细节学习率、精度与可视化模型结构定了之后能不能训好基本取决于训练配置。这一节讲的是我在实际项目里验证有效的参数组合和工程技巧不是照抄论文的默认值。5.1 优化器参数与学习率调度的实测组合ViT 微调我固定用 AdamW原因很简单它的权重衰减是对参数本身做的而不是像 Adam 那样把衰减混进梯度里对 Transformer 这类参数分布差异大的模型更友好。我的参数模板参数从头训练微调预训练模型optimizerAdamWAdamWbase lr3e-41e-5 ~ 5e-5weight decay0.050.01 ~ 0.05betas(0.9, 0.999)(0.9, 0.999)warmup epochs102 ~ 5schedulercosinecosine分层学习率不使用backbone 为 head 的 0.1 倍warmup 的作用是让学习率从很小的值线性爬升到 base lr。为什么必须加因为训练初期参数是随机的梯度方向噪声大直接用大学习率会让模型跑进一个很差的山谷。我一般在 warmup 阶段用LambdaLR实现def build_scheduler(optimizer, warmup_epochs, total_epochs, steps_per_epoch): warmup_steps warmup_epochs * steps_per_epoch total_steps total_epochs * steps_per_epoch def fn(step): if step warmup_steps: return step / max(1, warmup_steps) progress (step - warmup_steps) / max(1, total_steps - warmup_steps) return 0.5 * (1 math.cos(math.pi * progress)) return torch.optim.lr_scheduler.LambdaLR(optimizer, fn)steps_per_epoch千万别用len(dataloader)直接算如果你开了梯度累积真实的优化步数是 dataloader 长度除以累积步数算错了 cosine 曲线会提前触底。5.2 混合精度、梯度累积与数据加载的调优混合精度这块有几个细节值得说。autocast会自动把卷积和矩阵乘转成 fp16但 softmax、LayerNorm 这些对数值范围敏感的算子会保留 fp32这是它比手动.half()安全的原因。配合GradScaler处理梯度下溢scaler torch.cuda.amp.GradScaler() for step, (imgs, labels) in enumerate(loader): imgs, labels imgs.cuda(), labels.cuda() with torch.cuda.amp.autocast(): logits model(imgs) loss criterion(logits, labels) / accum_steps scaler.scale(loss).backward() if (step 1) % accum_steps 0: scaler.step(optimizer) scaler.update() scheduler.step() optimizer.zero_grad(set_to_noneTrue)注意loss / accum_steps这一步如果不除累积后的梯度会等比放大等效于把学习率乘了 accum_steps 倍。另外optimizer.zero_grad(set_to_noneTrue)比默认的置零更省内存因为它直接释放梯度张量而不是写零。DataLoader 这边num_workers设成 CPU 核心数的 0.7 倍左右比较合适pin_memoryTrue配合non_blockingTrue能进一步压缩数据搬运时间。我实测在单卡上做 224 分辨率训练数据加载经常是瓶颈尤其是开了 RandAugment 之后所以增强操作尽量放在 GPU 上或者用 DALI 之类的加速库。5.3 评估阶段别只看准确率混淆矩阵和注意力图更有信息量人脸情感识别是不平衡多分类任务只看 accuracy 会被多数类带偏。我固定输出三个东西macro F1每类 F1 的平均值对少数类敏感混淆矩阵一眼就能看出模型到底把哪两类搞混了。经验上恐惧和惊讶、悲伤和中性是最容易混的两组如果混淆矩阵里这两块的数字特别大说明模型在学纹理而不是在学表情结构。注意力图把最后一层某个 head 的注意力权重可视化出来叠加在原图上看模型有没有关注眼睛、嘴角这些关键区域。def visualize_attn(model, x, layer_idx-1, head_idx0): attentions [] hooks [] for blk in model.blocks: hooks.append(blk.attn.register_forward_hook( lambda m, i, o: attentions.append(o.detach()))) with torch.no_grad(): model(x) for h in hooks: h.remove() attn attentions[layer_idx][0, head_idx, 0, 1:] # cls 对所有 patch 的注意力 grid int(attn.shape[0] ** 0.5) attn attn.reshape(grid, grid).cpu() attn (attn - attn.min()) / (attn.max() - attn.min() 1e-8) return attn注意早点移除 hook否则多次调用会不断累积最后显存被吃光。这个坑我在一次多轮评估里踩过跑到第 30 个 batch 才 OOM查了半天。6. 一份可以抄作业的最小ViT实现前面讲了原理和坑这一节给一份能直接跑的完整实现。我刻意写得紧凑去掉了大部分锦上添花的模块方便你先跑通再按需扩展。6.1 模型定义的完整代码import math import torch import torch.nn as nn class DropPath(nn.Module): def __init__(self, p0.0): super().__init__() self.p p def forward(self, x): if self.p 0.0 or not self.training: return x keep 1 - self.p mask x.new_empty((x.shape[0],) (1,) * (x.ndim - 1)).bernoulli_(keep) return x * mask / keep class Mlp(nn.Module): def __init__(self, dim, hidden_dim, drop0.0): super().__init__() self.fc1 nn.Linear(dim, hidden_dim) self.act nn.GELU() self.fc2 nn.Linear(hidden_dim, dim) self.drop nn.Dropout(drop) def forward(self, x): return self.drop(self.fc2(self.drop(self.act(self.fc1(x))))) class Block(nn.Module): def __init__(self, dim, heads, mlp_ratio4.0, drop0.0, drop_path0.0): super().__init__() self.norm1 nn.LayerNorm(dim, eps1e-6) self.attn nn.MultiheadAttention(dim, heads, dropoutdrop, batch_firstTrue) self.norm2 nn.LayerNorm(dim, eps1e-6) self.mlp Mlp(dim, int(dim * mlp_ratio), drop) self.drop_path DropPath(drop_path) def forward(self, x): h self.norm1(x) h, _ self.attn(h, h, h, need_weightsFalse) x x self.drop_path(h) x x self.drop_path(self.mlp(self.norm2(x))) return x class MiniViT(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, num_classes7, embed_dim384, depth6, heads6, mlp_ratio4.0, drop0.1, drop_path0.1): super().__init__() self.grid img_size // patch_size self.num_patches self.grid ** 2 self.patch_embed nn.Conv2d(in_chans, embed_dim, patch_size, patch_size) self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed nn.Parameter(torch.zeros(1, self.num_patches 1, embed_dim)) self.pos_drop nn.Dropout(drop) dpr [drop_path * i / max(1, depth - 1) for i in range(depth)] self.blocks nn.ModuleList([ Block(embed_dim, heads, mlp_ratio, drop, dpr[i]) for i in range(depth) ]) self.norm nn.LayerNorm(embed_dim, eps1e-6) self.head nn.Linear(embed_dim, num_classes) nn.init.trunc_normal_(self.cls_token, std0.02) nn.init.trunc_normal_(self.pos_embed, std0.02) self.apply(self._init_weights) def _init_weights(self, m): if isinstance(m, nn.Linear): nn.init.trunc_normal_(m.weight, std0.02) if m.bias is not None: nn.init.zeros_(m.bias) elif isinstance(m, nn.LayerNorm): nn.init.zeros_(m.bias) nn.init.ones_(m.weight) def forward(self, x): B x.shape[0] x self.patch_embed(x).flatten(2).transpose(1, 2) # (B, N, C) cls self.cls_token.expand(B, -1, -1) x torch.cat([cls, x], dim1) # (B, N1, C) x self.pos_drop(x self.pos_embed) for blk in self.blocks: x blk(x) x self.norm(x) return self.head(x[:, 0])这份代码里有两个地方和官方实现不同我特意换成了更易懂的写法一是用nn.MultiheadAttention替代手写 QKV 投影二是用列表推导生成 drop path 率。如果你要复现论文精度还是建议换回手写注意力并加上qkv_bias。6.2 训练与推理脚本的骨架def train_one_epoch(model, loader, optimizer, scaler, scheduler, criterion, device, accum4): model.train() total, correct, loss_sum 0, 0, 0.0 for step, (imgs, labels) in enumerate(loader): imgs, labels imgs.to(device), labels.to(device) with torch.cuda.amp.autocast(): logits model(imgs) loss criterion(logits, labels) / accum scaler.scale(loss).backward() if (step 1) % accum 0: scaler.step(optimizer) scaler.update() scheduler.step() optimizer.zero_grad(set_to_noneTrue) loss_sum loss.item() * accum correct (logits.argmax(1) labels).sum().item() total labels.size(0) return loss_sum / total, correct / total torch.no_grad() def evaluate(model, loader, device): model.eval() correct, total 0, 0 for imgs, labels in loader: imgs, labels imgs.to(device), labels.to(device) pred model(imgs).argmax(1) correct (pred labels).sum().item() total labels.size(0) return correct / total # 组装训练流程 model MiniViT(num_classes7).to(cuda) criterion nn.CrossEntropyLoss(label_smoothing0.1) optimizer torch.optim.AdamW(model.parameters(), lr3e-4, weight_decay0.05) scheduler build_scheduler(optimizer, warmup_epochs5, total_epochs50, steps_per_epochlen(train_loader) // 4) scaler torch.cuda.amp.GradScaler() for epoch in range(50): tr_loss, tr_acc train_one_epoch(model, train_loader, optimizer, scaler, scheduler, criterion, cuda) val_acc evaluate(model, val_loader, cuda) print(fepoch {epoch:03d} | loss {tr_loss:.4f} | train {tr_acc:.4f} | val {val_acc:.4f}) torch.save(model.state_dict(), fvit_epoch{epoch}.pth)6.3 跑通之后可以继续深挖的方向代码跑起来之后有几个方向值得继续折腾。第一是替换注意力实现用F.scaled_dot_product_attention替代手写矩阵乘在支持 flash attention 的显卡上速度能提升一倍以上。第二是加蒸馏用 EfficientNetV2 当教师网络去指导 ViT 训练这个组合在小数据集上经常比单独训 ViT 高出两三个点。第三是换更细的 patch 划分比如 patch_size8token 数变成 784注意力计算量翻四倍但细粒度表情的细节保留更完整适合算力充裕的场景。最后一个我在实际调试中养成的习惯每次改动模型结构后先用一个 batch 的数据跑一遍完整的训练步确认 loss 能回传、参数确实更新了再投入长周期训练。这一步看起来多余但它帮我抓出过好几次梯度被 detach 掉了某层没进优化器的静默错误——这类问题不会报错只会让你在训练十个小时之后发现 loss 是一条直线。
返回列表