实现类BERT模型完成MNIST分类的技术疑问与代码问题
任务背景与实现细节
我正在自学并实现一个类BERT网络,用于MNIST数据集的CV特征提取与分类任务。该任务简单适合学习,尽管Transformer编码器并非图像特征提取的最优方案,但足以验证效果。
我的实现细节如下:
- 使用torchvision提供的MNIST数据集,batch_size设为128时,dataloader返回形状为
[128,1,28,28]的28*28单通道图像; - 将每张图像切割为16个均等块(横竖各4份),将图像数据转为长度为16、维度为
(28/4)²=49的序列,即[128,16,49]; - 通过全连接层转换为隐藏维度,得到
[128,16,64]; - 重复6个结构类似BERT层的Transformer编码器层;
- 通过MLP输出最终分类结果,即
[128,10]。
细节疑问解答
1. MNIST预处理的原因
示例代码中对MNIST数据集做如下预处理:
transform_mnist = torchvision.transforms.Compose([ torchvision.transforms.ToTensor(), torchvision.transforms.Normalize((0.1307,), (0.3081,)) ])
这和位置嵌入的分布要求无关,核心是标准化数据分布:
ToTensor()将图像从0-255的uint8格式转为0-1的float32格式;Normalize用MNIST数据集的全局均值(0.1307)和标准差(0.3081)做标准化,让数据分布接近均值为0、方差为1的正态分布,能加速模型收敛,避免输入数据分布差异导致的训练不稳定。
2. cls_token的作用
cls_token将原始数据从[128,16,64]转为[128,17,64],核心作用有两点:
- 作为全局特征聚合载体:通过自注意力机制与所有图像块特征交互,最终融合整个图像的全局信息;
- 专门服务分类任务:后续MLP直接使用cls_token的输出做分类,避免从所有图像块中额外聚合特征的操作,简化任务流程。
3. 多头自注意力中的Dropout使用
生产环境中,BERT及大部分变体都会在多头自注意力中使用Dropout:
- 通常在注意力权重计算后(softmax之后)添加Dropout,防止模型过拟合;
- 你的代码在
W_o之后加Dropout是合理的;部分实现也会在注意力分数上直接加Dropout,两种方式都被广泛使用,核心都是正则化模型。
4. LayerNorm、Dropout与残差连接的顺序
BERT标准结构采用Pre-LN(LayerNorm放在子模块前),也有Post-LN的变体,推荐顺序分两种:
- 标准BERT结构(Pre-LN):
LayerNorm -> 自注意力/FFN -> Dropout -> 残差连接,这种结构训练更稳定,梯度消失风险更低; - Post-LN(早期Transformer结构):
自注意力/FFN -> Dropout -> 残差连接 -> LayerNorm,你的代码用的是Post-LN,但EncoderBlock的forward方法中最后误用了norm_layer1,应该改为norm_layer2,属于代码错误。
5. 单次位置编码的有效性
示例代码仅执行一次位置编码是合理的,后续块依然能捕捉位置信息:
- 位置编码在输入时添加到特征中,后续每个编码器块的自注意力机制会基于包含位置信息的特征计算,位置信息会在整个编码过程中传递;
- 不需要每个块都加位置编码,否则会导致位置信息重复叠加,破坏原始特征分布。
6. MLP输入选择的原因
最终MLP使用x[:,0,:](即cls_token的输出)是因为:
- 分类任务需要全局特征,cls_token经过所有自注意力层交互后,已经融合了整个图像的信息,适合作为分类输入;
- 回归任务中,如果是全局回归(比如预测图像的某个全局属性),同样可以使用cls_token的输出;如果是像素级回归,则需要使用所有图像块的特征做后续处理。
代码训练异常排查(输出和损失持续累积)
你的代码可运行但损失持续累积,核心问题有以下几点:
1. 损失函数使用错误
F.nll_loss要求输入是对数概率,但你的模型最后一层直接输出分类logits,未经过log_softmax处理,导致损失计算错误。推荐修改训练循环中的损失计算:
loss = F.cross_entropy(output, y)
cross_entropy会自动包含softmax和nll_loss的计算,更适配当前模型输出。
2. EncoderBlock的LayerNorm误用
在EncoderBlock的forward方法中,最后返回时错误使用了norm_layer1,应该改为norm_layer2:
def forward(self, x): y = self.self_attention_layer(x) z = self.norm_layer1(x + y) return self.norm_layer2(z + self.ffn(z)) # 将norm_layer1改为norm_layer2
3. 学习率过高
当前设置的lr=0.002对于Transformer模型来说偏高,容易导致训练不稳定、损失震荡或持续上升。建议调整为lr=1e-4或5e-4,也可配合学习率调度器(如torch.optim.lr_scheduler.StepLR)优化学习过程。
4. 缺少训练模式切换
建议在训练循环开头添加net.train(),确保模型处于训练模式(Dropout等正则化层生效):
for epoch in range(100): net.train() for step, (x, y) in enumerate(train_loader): # 现有训练代码
内容的提问来源于stack exchange,提问作者AdamHommer

