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

使用nn.TransformerEncoder训练ViT分类模型时损失不下降的问题

ViT图像分类中PyTorch原生TransformerEncoder训练损失不下降问题

问题现象

  • 使用第三方实现的another_transformer_encoder时,模型训练正常,损失可正常下降
  • 使用PyTorch原生nn.TransformerEncoder(对应代码中的transformer_encoder)时,模型损失无法正常下降
  • 已在TransformerEncoderLayer中设置batch_first=True,但问题未解决

模型代码

import torch
import torch.nn as nn
from torch.nn import TransformerEncoderLayer, TransformerEncoder

# 自定义位置编码实现
class PositionalEncoding(nn.Module):
    def __init__(self, dim, max_len=5000):
        super().__init__()
        position = torch.arange(max_len).unsqueeze(1)
        div_term = torch.exp(torch.arange(0, dim, 2) * (-torch.log(torch.tensor(10000.0)) / dim))
        pe = torch.zeros(max_len, 1, dim)
        pe[:, 0, 0::2] = torch.sin(position * div_term)
        pe[:, 0, 1::2] = torch.cos(position * div_term)
        self.register_buffer('pe', pe)

    def forward(self, x):
        x = x + self.pe[:x.size(1)].transpose(0, 1)
        return x

# 第三方Transformer示例实现(仅示意)
class Transformer(nn.Module):
    def __init__(self, dim, depth, heads, mlp_dim):
        super().__init__()
        self.layers = nn.ModuleList([])
        for _ in range(depth):
            self.layers.append(nn.ModuleList([
                nn.LayerNorm(dim),
                nn.MultiheadAttention(dim, heads, batch_first=True),
                nn.LayerNorm(dim),
                nn.Sequential(
                    nn.Linear(dim, mlp_dim),
                    nn.GELU(),
                    nn.Linear(mlp_dim, dim),
                    nn.Dropout(0.5)
                )
            ]))
    
    def forward(self, x):
        for norm1, attn, norm2, mlp in self.layers:
            x = x + attn(norm1(x), norm1(x), norm1(x))[0]
            x = x + mlp(norm2(x))
        return x

class ViT(nn.Module):
    def __init__(self, image_size, patch_size, channels, num_classes, dim, depth, heads, mlp_dim):
        super(ViT, self).__init__()

        # 计算patch数量
        self.num_patches = (image_size // patch_size) ** 2

        # Patch嵌入层
        self.patch_embedding = nn.Conv2d(channels, dim, kernel_size=patch_size, stride=patch_size, bias=False)

        # PyTorch原生TransformerEncoder
        encoder_layers = TransformerEncoderLayer(
            d_model=dim,
            nhead=heads,
            dim_feedforward=mlp_dim,
            dropout=0.5,
            batch_first=True
        )
        self.transformer_encoder = TransformerEncoder(
            encoder_layer = encoder_layers,
            num_layers=depth
        )
        # 第三方Transformer实现
        self.another_transformer_encoder = Transformer(dim, depth, heads, mlp_dim)

        # 位置编码层
        self.pos_encoder = PositionalEncoding(dim)

        # 分类头
        self.cls = nn.Parameter(torch.randn(1, 1, dim))
        self.classification_head = nn.Linear(dim, num_classes)

    def forward(self, x):
        # 图像分patch
        patches = self.patch_embedding(x)    # (batch_size, dim, num_patches_h, num_patches_w)
        patches = patches.flatten(2)   # (batch_size, dim, num_patches)
        patches = patches.transpose(1, 2)    # (batch_size, num_patches, dim)

        # 添加位置编码
        patches = self.pos_encoder(patches)

        # 添加CLS token
        cls_token = self.cls.expand(x.shape[0], -1, -1)    # (batch_size, 1, dim)
        patches = torch.cat([cls_token, patches], dim=1)   # (batch_size, num_patches+1, dim)

        # 送入Transformer编码器
        patches = self.transformer_encoder(patches) # 存在问题的分支
        # patches = self.another_transformer_encoder(patches) # 正常运行的分支

        # 提取CLS token做分类
        cls_token = patches[:, 0]
        output = self.classification_head(cls_token)
        return output

模型调用代码

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = ViT(image_size=28,
            patch_size=7,
            channels=1,
            num_classes=10,
            dim=64,
            depth=6,
            heads=8,
            mlp_dim=128
            ).to(device)
criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.SGD(model.parameters(), lr=0.1, momentum=0.9)

问题排查与解决方案

1. 层归一化位置不匹配

PyTorch原生TransformerEncoderLayer默认使用Pre-LN(在多头注意力和MLP模块输入前做归一化),而绝大多数ViT第三方实现采用Post-LN(在模块输出后做归一化)。Pre-LN在高学习率下更容易出现训练不稳定,导致损失无法下降。

解决方法:修改TransformerEncoderLayer的norm_first参数为False,切换为Post-LN:

encoder_layers = TransformerEncoderLayer(
    d_model=dim,
    nhead=heads,
    dim_feedforward=mlp_dim,
    dropout=0.5,
    batch_first=True,
    norm_first=False  # 关键修改:使用Post-LN
)

2. 位置编码格式验证

确保自定义PositionalEncoding的输出维度与batch_first=True的输入格式匹配。原生Transformer在batch_first=True时期望输入为(batch_size, seq_len, dim),位置编码需对应调整维度。

验证点:PositionalEncoding的forward方法中,需将位置编码的维度从(seq_len, 1, dim)转置为(1, seq_len, dim),再与输入相加(输入为(batch, seq_len, dim))。

3. 优化器与学习率调整

原生Transformer对学习率的敏感度高于第三方ViT实现,当前使用的SGD+0.1学习率可能过高。

推荐调整:

  • 改用ViT常用的AdamW优化器,并降低学习率:
optimizer = torch.optim.AdamW(model.parameters(), lr=0.001, weight_decay=0.01)
  • 若坚持使用SGD,可将学习率降至0.01或更低,并添加权重衰减。

4. 参数初始化优化

第三方Transformer通常会对注意力权重、层归一化参数做针对性初始化,而原生TransformerEncoder的默认初始化可能不够适配ViT任务。

手动初始化示例:

def __init__(self, ...):
    # ... 其他初始化代码 ...
    # 初始化TransformerEncoder参数
    for p in self.transformer_encoder.parameters():
        if p.dim() > 1:
            nn.init.xavier_uniform_(p)  # 用Xavier初始化权重
        else:
            nn.init.zeros_(p)  # 偏置初始化为0

内容的提问来源于stack exchange,提问作者Dylan Chiu

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 14:18:16