You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

实现类BERT模型完成MNIST分类的技术疑问与代码问题

类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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.24 19:18:19