基于Pytorch复现Vision Transformer:原理剖析与完整实现
发布时间:2026/9/29 8:47:30来源:尧图网络
Vision TransformerViT这篇论文在2020年出来的时候业界最大的争议并不是“Transformer能不能做视觉”而是“为什么这个看起来如此简单、几乎不带任何视觉先验的模型能在数据量足够大的情况下赢过当时调参到极致的大规模CNN”。我自己的复现动机也很朴素想把Pytorch版的ViT从Patch Embedding到最后的分类头逐行写出来搞懂而不是直接import vit_pytorch了事。这个系列记录的就是我在复现ViT模型时踩过的坑、验证过的推导以及一份可以直接copy下来跑的完整实现。如果你对Vision Transformer在视觉上的原理还有模糊的地方或者有一坨代码不知道从哪里下笔这篇笔记应该能帮你省下不少调试时间。1. 为什么ViT敢把卷积扔掉设计动机与整体管线1.1 CNN的建筑先验和Transformer的“平权”思想要复现一个模型第一步不是抄代码而是搞清楚作者当时在想什么。CNN之所以在图像领域统治了这么多年靠的是三个内置先验局部连接卷积核只看小邻域、权值共享同一个卷积核扫过全图和平移等变性目标换个位置特征也跟着平移。这三个先验对自然图像非常有效因为它们假设了“相邻像素相关性高、远处的像素关系可以靠堆层数慢慢建立”。但Transformer的出发点完全不同。自注意力机制的一个关键特性是它在一开始就让序列里的任意两个位置直接交互不经过中间层的“传递”。把图像拉成一串token之后第1层就能建立远距离依赖这对理解全局语义、物体之间的相对关系很有帮助。代价是它几乎没有视觉先验所有关于“图像到底是什么结构”的知识都得靠数据自己学出来。这解释了ViT论文里一个核心结论模型的结构性先验越强越省数据先验越弱越依赖数据规模。ViT在ImageNet-1k上用从头训练的方式和同规模CNN比是有差距的但在ImageNet-21k甚至JFT-300M这种大规模数据集上预训练之后再去下游任务微调就能反超同期CNN。复现这个模型之前脑子里得有这根弦如果你手里的数据很少直接用ViT裸训大概率会过拟合这不是代码问题是模型本身的设计取向问题。1.2 ViT的整体数据流ViT-Base这个最常复现的版本参数大约86M左右核心配置是embed_dim768, depth12, num_heads12, mlp_ratio4。它的整体流程可以这样理解一张224x224的图先切成14x14个16x16的patch每个patch经过线性投影变成768维向量于是图像变成了长度为196的“句子”。句子开头再拼接一个用于分类的cls_token加上位置编码后送入12层Transformer Encoder。最后取出cls_token对应的输出过一个LayerNorm和一个线性分类头得到类别预测。整个流程里最反直觉的一点是图像的空间结构几乎只在Patch Embedding和位置编码里体现了一次之后Encoder本身对“图像”一无所知。它看到的纯粹是一个序列。这个设计说好听是简洁说难听是浪费了大量先验知识但事实证明大规模预训练确实能让模型自己学会空间关系。所以复现时也不要想着给它加各种花哨的视觉先验先把论文这套“极简方案”跑通再去想怎么改。2. 复现前的准备工作环境、数据与工程骨架2.1 环境配置复现ViT不需要特别重的依赖。我本地的环境是Pytorch 2.x、CUDA 11.8、torchvision 0.17Python 3.10。如果你是CPU环境也能跑通前向和反向只是训练会非常慢所以建议尽量准备一张NVIDIA显卡。另外强烈建议安装timm倒不是要用它把模型直接实例化而是因为社区里大量的训练配置、预训练权重都围绕timm的接口展开复现过程中对照它的源码能省很多事。依赖清单大致是pip install torch torchvision pip install timm einops tqdm tensorboardeinops不是必需的但用rearrange写维度变换时非常直观尤其适合在Attention里处理多头拆分。tensorboard用于观察loss曲线和attention map是我调试时的标配。2.2 数据选择与工程目录ViT的论文实验是在ImageNet上做的224x224的输入、300个epoch、batch size 4096单卡复现这个配置不现实。我自己复现时的路径是先把模型结构和单卡可完成的训练闭环在CIFAR-10上跑通再切换到大一点的输入或数据集上验证。CIFAR-10的图是32x32相比224小很多用patch_size4、embed_dim256、depth6这种微型变体半小时到一小时就能看到一个明确的收敛趋势非常适合验证代码有没有写错。工程目录我建议这样组织vit_impl/ ├── vit_model.py # 模型定义 ├── train.py # 训练脚本 ├── data.py # 数据加载与增强 ├── config.py # 超参配置 └── utils.py # 工具函数种子设置、可视化等一个容易被忽略的点是固定随机种子。Transformer的初始化方差比较大如果不固定种子同一个模型两次训练可能得到完全不同的结果排查问题时很难判断是代码问题还是随机性导致的。我一般在train.py入口做这几件事import random import numpy as np import torch def set_seed(seed42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed)这样后续调整超参时至少能把变量控制在模型本身。3. Patch Embedding与位置编码图像如何“变成”一句话3.1 为什么用Conv2d实现Patch切分ViT坐标变换的第一步是把(B, 3, 224, 224)的图变成(B, 196, 768)的token序列。论文里说的是“把每个16x16的patch线性投影到768维”初学者容易把它拆成两步先切割patch再对每个patch拉直后接Linear。实际上在实现时我用了一个nn.Conv2d一步搞定效果完全等价。import torch import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, embed_dim768): super().__init__() self.img_size img_size self.patch_size patch_size self.num_patches (img_size // patch_size) ** 2 self.proj nn.Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): B, C, H, W x.shape assert H self.img_size and W self.img_size, \ f输入尺寸应为 {self.img_size}x{self.img_size} x self.proj(x) # (B, embed_dim, H/P, W/P) x x.flatten(2) # (B, embed_dim, num_patches) x x.transpose(1, 2) # (B, num_patches, embed_dim) return x为什么Conv2d能做这件事因为卷积本质上就是滑窗内像素的加权求和。当kernel_size16, stride16时每个窗口覆盖一个不重叠的16x16 patch输出通道数等于embed_dim每个输出通道就相当于把3*16*16768个输入像素做一次线性组合。不同patch用的是同一组卷积核这正好对应了Linear里“权重共享”的含义。所以从数学上讲Conv2d在这里和“把patch拉平再过一个Linear”是严格等价的但Conv2d在GPU上更高效也省去了切割patch的麻烦。这里有一个细节值得注意如果你的embed_dim恰好等于patch_size*patch_size*in_chans比如输入是单通道灰度图、patch_size16、embed_dim256那Conv2d的每个输出通道其实就是对某个像素位置的“直接采样”相当于没学权重。实际项目里embed_dim通常大于这个值所以不用太担心。3.2 class token与位置编码图像变成token序列后还需要两个准备工作拼上cls_token加上位置编码。cls_token是从BERT里继承过来的习惯它的作用是提供一个“全局汇聚”的位置让模型在最后一层只用看这个token的输出就能做分类而不是去池化所有patch token。它是可学习参数初始化为0或很小的随机数。class_token nn.Parameter(torch.zeros(1, 1, embed_dim)) position_embedding nn.Parameter(torch.zeros(1, num_patches 1, embed_dim))这里的num_patches 1是196个patch加上1个cls token。位置编码直接用1D可学习向量即可。论文里比较过1D、2D和相对位置编码结论是在大规模预训练下差别不大所以复现代码基本都用1D简洁且容易实现。前向时把cls token批量复制并拼接在序列前面x patch_embed(x) # (B, N, C) cls_tokens class_token.expand(B, -1, -1) x torch.cat((cls_tokens, x), dim1) # (B, N1, C) x x position_embedding x pos_drop(x)pos_drop是一个dropout层默认概率0.1。位置编码是加在patch token和cls token上的不是拼接这一点初学者容易搞混。3.3 一个可复用的调试小技巧验证PatchEmbed的输出形状每次刚写完PatchEmbed我建议都单独跑一个前向断言确认形状符合预期再看下一步。比如model PatchEmbed(img_size224, patch_size16, embed_dim768) x torch.randn(2, 3, 224, 224) out model(x) print(out.shape) # 期望 torch.Size([2, 196, 768])这种“模块级验证”能帮你把维度问题扼杀在早期而不是等整个模型堆完再排查。我实际写模型时每完成一个子模块就打印一次形状几乎已经成了肌肉记忆。4. Encoder核心Attention、MLP与残差连接的Pytorch实现4.1 Attention模块的完整理解Transformer Encoder里最核心的部分就是多头自注意力。ViT沿用了NLP Transformer的原始设计只是输入换成了图像token序列。它的思想不复杂对序列里的每个token计算它和其他所有token的相关性再用相关性加权聚合其他token的信息。实现时一个常用的优化是用一次Linear同时生成Q、K、V而不是写三个独立的Linear。这样参数量一样但计算更紧凑也是timm源码的风格。看这段代码时要特别注意维度变换class Attention(nn.Module): def __init__(self, dim, num_heads8, qkv_biasFalse, attn_drop0., proj_drop0.): super().__init__() self.num_heads num_heads self.scale (dim // num_heads) ** -0.5 self.qkv nn.Linear(dim, dim * 3, biasqkv_bias) self.attn_drop nn.Dropout(attn_drop) self.proj nn.Linear(dim, dim) self.proj_drop nn.Dropout(proj_drop) def forward(self, x): B, N, C x.shape qkv self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads) qkv qkv.permute(2, 0, 3, 1, 4) q, k, v qkv.unbind(0) # 每个形状 (B, num_heads, N, head_dim) attn (q k.transpose(-2, -1)) * self.scale attn attn.softmax(dim-1) attn self.attn_drop(attn) x attn v # (B, num_heads, N, head_dim) x x.transpose(1, 2).reshape(B, N, C) x self.proj(x) x self.proj_drop(x) return x, attn这里绕的地方在于reshape和permute的顺序。reshape(B, N, 3, num_heads, head_dim)先显式把最后一维分成四段然后permute(2,0,3,1,4)把“3”这个维度挪到最前面这样unbind(0)解出来正好是q、k、v三个张量。如果你觉得这个写法不好读也可以写成三个独立的Linear分别生成q、k、v效果理论上一致只是少了一点参数共享和计算效率。我建议至少在初学阶段用一次Linear生成QKV的写法因为这是工业界最常见的实现很多预训练权重也是按这个结构存参数的。多头注意力的意义在于让模型同时从多个子空间观察序列关系。每个头有自己的Q、K、V投影注意力权重可以关注不同的pattern比如有的头关注相邻patch有的头关注全局语义区域。这也是后面做attention map可视化时的观察对象。4.2 MLP与Block结构Encoder的每个Block由两层子结构组成Attention和MLP各自外面都套了LayerNorm和残差连接。MLP的实现比较简单但有一个细节要注意隐藏层维度是输入维度的4倍也就是768 - 3072 - 768中间用GELU激活。class Mlp(nn.Module): def __init__(self, in_features, hidden_featuresNone, out_featuresNone, act_layernn.GELU, drop0.): super().__init__() hidden_features hidden_features or in_features out_features out_features or in_features self.fc1 nn.Linear(in_features, hidden_features) self.act act_layer() self.fc2 nn.Linear(hidden_features, out_features) self.drop nn.Dropout(drop) def forward(self, x): x self.fc1(x) x self.act(x) x self.drop(x) x self.fc2(x) x self.drop(x) return x接下来把它组装成Block。很多初学者会先写Attention再写MLP然后把残差加在外面但要注意归一化的位置。class Block(nn.Module): def __init__(self, dim, num_heads, mlp_ratio4., drop0., attn_drop0.): super().__init__() self.norm1 nn.LayerNorm(dim) self.attn Attention(dim, num_headsnum_heads, qkv_biasTrue, attn_dropattn_drop, proj_dropdrop) self.norm2 nn.LayerNorm(dim) self.mlp Mlp(in_featuresdim, hidden_featuresint(dim * mlp_ratio), act_layernn.GELU, dropdrop) def forward(self, x): x x self.attn(self.norm1(x))[0] x x self.mlp(self.norm2(x)) return x4.3 为什么ViT用Pre-LN而不是Post-LN这里值得单独说一嘴。原始的Transformer论文用的是Post-LN也就是在残差相加之后再过LayerNorm即x LN(x Attn(x))。而ViT用的是Pre-LN也就是先归一化再进入Attention/MLP最后做残差相加即x x Attn(LN(x))。Pre-LN的好处是梯度可以直接从输出端流回输入端基本不受层数影响训练更稳定可以允许相对较大的学习率。ViT的层数有12层甚至24层用Post-LN在训练初期很容易出现震荡。这也是为什么你在timm或官方实现里看到的ViT都是Pre-LN。这里提醒一点结构图上画的“Add Norm”有两种摆放方式复现时一定要看清楚是Norm在Add之前还是之后直接抄代码很容易抄错。Block里Attention返回的是(x, attn)前向时用[0]取主输出。这个设计是为了方便后续可视化attention map如果不需要可视化也可以只返回x。5. 组装VisionTransformer主类与参数量核对5.1 VisionTransformer主类实现有了PatchEmbed、Block、Attention、Mlp这些积木主类就非常清晰了。整体结构是PatchEmbed - 拼接cls token - 加位置编码 - 进12层Block - 最后一层LayerNorm - 取出cls token输出 - 分类head。class VisionTransformer(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, num_classes1000, embed_dim768, depth12, num_heads12, mlp_ratio4., drop_rate0., attn_drop_rate0.): super().__init__() self.num_classes num_classes self.num_features embed_dim self.patch_embed PatchEmbed(img_sizeimg_size, patch_sizepatch_size, in_chansin_chans, embed_dimembed_dim) num_patches self.patch_embed.num_patches self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed nn.Parameter(torch.zeros(1, num_patches 1, embed_dim)) self.pos_drop nn.Dropout(pdrop_rate) self.blocks nn.Sequential(*[ Block(dimembed_dim, num_headsnum_heads, mlp_ratiomlp_ratio, dropdrop_rate, attn_dropattn_drop_rate) 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): nn.init.trunc_normal_(self.pos_embed, std0.02) nn.init.trunc_normal_(self.cls_token, std0.02) self.apply(self._init_module) def _init_module(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.ones_(m.weight) nn.init.zeros_(m.bias) def forward(self, x): B x.shape[0] x self.patch_embed(x) cls_tokens self.cls_token.expand(B, -1, -1) x torch.cat((cls_tokens, x), dim1) x self.pos_drop(x self.pos_embed) x self.blocks(x) x self.norm(x) x self.head(x[:, 0]) return x前向过程里有个容易混淆的点x[:, 0]取的是序列维度上的第0个位置也就是cls token的输出。它不是取batch的第0个样本也不是取特征维的某个通道。如果你对张量索引不熟建议在这行加个print(x.shape)看一下维度。5.2 参数初始化的细节初始化看起来琐碎但对训练稳定性影响很大。ViT官方和timm里pos_embed和cls_token都用trunc_normal_(std0.02)初始化Linear权重同样用trunc_normalLayerNorm的weight初始化成1、bias初始化成0。之所以用trunc_normal而不是普通正态分布是因为它能截断过大值避免某些token在初始化阶段就产生过大的激活值从而破坏后续LayerNorm的数值范围。如果你是从零开始复现用Pytorch默认的kaiming_uniform_或默认初始化通常也能训但loss曲线可能没有用trunc_normal那么平稳。既然复现的目标是尽量贴近原论文行为建议还是把初始化逻辑抄全。5.3 参数量快速核算写完模型后一个非常有效的验证方法是和公开权重核对参数量。标准ViT-B/16输入224、patch16、embed_dim768、depth12的参数量大约在86M级别。手工粗算也很有意思12层Block里每层包含Attention和MLP。Attention中qkv线性是768*768*3约1.77Mproj是768*768约0.59M加起来每层约2.36M12层约28.3M。MLP中fc1是768*3072约2.36Mfc2相同每层约4.72M12层约56.6M。PatchEmbed的卷积核是3*768*16*16约0.59M。位置编码197*768约0.15M分类头768*1000约0.77M。加总大概在86.4M左右和公开的86M量级对得上。如果你打印出来的参数量和这个差太远说明某个模块写错了比如mlp_ratio没有乘进去或者num_heads设置不对。model VisionTransformer(img_size224, patch_size16, embed_dim768, depth12, num_heads12, num_classes1000) total_params sum(p.numel() for p in model.parameters()) print(fTotal params: {total_params / 1e6:.2f}M)这一步建议在写完整模型后立刻做比自己反复读代码找bug快得多。6. 训练配置与超参经验让模型真正收敛6.1 优化器、学习率与warmup模型本身写完了只是迈出了第一步。ViT对训练超参比CNN敏感得多照搬CNN那套SGD加大学习率的方案会翻车。我自己跑下来比较稳的配置是AdamW优化器、lr3e-4、weight_decay0.05、使用warmup cosine衰减。ViT原论文用的是Adam而不是AdamW但社区主流复现都用AdamW效果更好也更常见。学习率的选择和batch size有关。经验上ViT在batch size 256到512之间lr从3e-4到1e-3都能接受。我习惯用一个简单公式估算初始lrlr base_lr * batch_size / 256比如batch size 128就用3e-4 * 128 / 256 1.5e-4左右起步。warmup建议至少设置总步数的5%我自己在CIFAR-10上通常warmup 500步左右让优化器的动量状态先稳定下来再进入正常学习阶段。还要提一个很实用的小技巧训练全程开启混合精度。ViT在batch size有限时混合精度能把显存占用降下来30%到40%训练速度也明显提升。Pytorch 2.x里用torch.autocast包住前向和loss计算再加GradScaler缩放梯度代码量不大收益明显。6.2 数据增强与小数据集应对如果只在CIFAR-10这种小数据集上从头训练ViT不加增强几乎是必过拟合的。我在复现时最少会加三种增强随机水平翻转、随机裁剪、以及一定程度上的颜色抖动。如果想让结果更好看可以再上RandAugment、Mixup和CutMix。但这三种属于“高阶操作”建议先在无增强的朴素配置下跑通整个链路再加上去调参。这里有一个来自论文的提醒ViT在ImageNet-1k上从头训练时如果只做基础增强效果会明显差于同级别CNN但配合强增强RandAugment、Mixup、CutMix和更长的训练周期差距会显著缩小。DeiT那篇论文专门研究过这件事结论是“Transformer在中小规模数据上训练数据增强不是可选项是必需品”。所以你在小数据集上看到ViT的loss下降慢、验证集上不去时先别急着怀疑代码优先排查数据增强强不强。6.3 训练循环骨架与一个小数据集的实测结论训练脚本的基本骨架不复杂这里我给一个精简版的伪码optimizer torch.optim.AdamW(model.parameters(), lrlr, weight_decay0.05) scheduler CosineAnnealingLR(optimizer, T_maxepochs) scaler torch.cuda.amp.GradScaler() for epoch in range(epochs): model.train() for images, labels in train_loader: images, labels images.cuda(), labels.cuda() optimizer.zero_grad() with torch.autocast(device_typecuda, dtypetorch.float16): outputs model(images) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() scheduler.step()我在CIFAR-10上的一个参考配置是输入分辨率32x32、patch_size4、embed_dim256、depth6、num_heads8、batch_size128、epochs100、用上面提到的增强组合。这种情况下验证集准确率通常能稳定到80%到85%之间。这个数字本身并不重要重要的是它告诉你一个信号如果训练了二三十个epoch验证集准确率还停留在三十以下那大概率不是模型表达能力的问题而是学习率、初始化或梯度链路出了问题需要回头排查。另外提一句ViT里没有BatchNorm全部是LayerNorm这意味着单卡batch size再小也能稳定训练不至于像某些带BN的模型那样batch太小导致统计量漂移。这算是它的一个隐性优点。7. 复盘复现过程踩过的坑与解决办法7.1 Loss不降反升或直接变NaN最常见的翻车点。我的排查顺序是这样的先用很小的lr比如1e-5跑几十步看loss有没有下降趋势如果下降说明问题出在lr或warmup配置上。ViT在训练初期对lr非常敏感尤其是batch size较小的场景lr稍微偏大就会出现loss不降甚至窜升的现象。另一个常见原因是qkv_bias没有打开导致Attention在初始化阶段的偏置为零训练早期梯度不稳定。timm默认把qkv_biasTrue建议照着设。如果loss直接变NaN优先检查数据里有没有异常值以及是否启用了混合精度但没有使用GradScaler。某些情况下mixup或cutmix产生的标签是平滑浮点数如果loss实现里写错了维度也容易出现数值问题。7.2 维度配不上ViT的维度变换集中在Attention里报错信息往往指向query和key的形状不匹配。这类问题最好的解决方式是在每个模块的输出处打印shape。我一般会在PatchEmbed、cat(cls_token)、加pos_embed、进入第一层Block之后分别打印一次。比如torch.Size([2, 197, 768])出现后后面所有Block的输入输出都应该保持这个形状。一旦某个Block内部改变了序列长度说明Attention的reshape写错了。还有一个隐蔽问题位置编码的长度num_patches 1如果你用torch.zeros(1, num_patches, embed_dim)忘了加1运行时会直接报错因为加完cls token后序列长度是197但pos_embed只有196。这个报错信息一般很明确一看就知道。7.3 Attention map可视化出来一片模糊训练结束后把某一层的attention拿出来画热力图是验证模型有没有学到空间关系的直观手段。我见过不少人的可视化图是一整块均匀颜色这通常有两种可能一是你取的是softmax之前的logits二是模型根本没有收敛。正确的做法是取Attention模块里softmax之后的权重形状是(B, num_heads, N, N)。想可视化cls token关注哪些patch就取attn[0, head_idx, 0, 1:]这对应cls token对所有patch token的注意力权重然后reshape成(14, 14)插值放大到和原图一样大叠加在输入图上。ViT浅层attention通常有一定局部性但不会像CNN那样只盯着很小邻域深层head会更关注语义上相关的区域。如果你看到的attention在整张图上均匀分布先确认模型在验证集上是否已经有一定准确率再检查可视化代码有没有取错张量。7.4 加载预训练权重时的键名匹配复现到后期难免想加载官方ViT-B/16权重来做微调或下游任务对比。这时最容易踩的坑是键名对不上。timm或官方仓库里权重文件的键名通常是patch_embed.proj.weight、blocks.0.attn.qkv.weight这类写法。如果你的模型里某个模块名字和官方不一样load_state_dict会直接报missing keys或unexpected keys。我的处理方法是先加载官方的state_dict打印出所有key然后和自己的模型对照。另外官方权重是在ImageNet-21k或ImageNet-1k上预训练的分类头是1000类。如果你要微调到10类的小数据集加载权重时一定要把分类头排除在外使用strictFalse并忽略head.weight、head.bias。否则模型的输出维度都不一致直接加载必然报错。state_dict torch.load(vit_base_patch16_224.pth, map_locationcpu) # 去掉分类头相关键 state_dict {k: v for k, v in state_dict.items() if not k.startswith(head.)} model.load_state_dict(state_dict, strictFalse)这一套处理同样适用于把ViT当成骨干网络接其他任务头的场景。我自己在复现到下游分割、检测任务时基本都是保留head之外的全部权重只替换任务头。7.5 复现时的一点心态建议ViT这类模型的结构复杂度不在“层数深”而在“模块多且相互依赖”。我踩坑后的体会是不要一上来就追求跑出顶尖精度先把代码跑通、参数量对上、小数据集上能收敛这三个里程碑完成后再逐步优化。特别是位置编码、Pre-LN、QKV投影这三处细节几乎是所有复现bug的高发区写的时候多留几个print(shape)多对照原论文的结构图能省下一晚上的调试时间。另外我强烈建议在复现完ViT之后顺手把timm的vit_base_patch16_224实例化成同样配置对比两者的前向输出。虽然随机初始化权重不同输出数值不会一致但只要形状、参数量、loss变化趋势都在合理范围基本可以认为复现是成功的。这个“对照实验”比盲目改代码有效得多。
网站建设高端定制企业官网