使用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
相关产品推荐
相关产品推荐

