单头Transformer模型输出与目标张量尺寸不匹配问题求助
单头Transformer尺寸不匹配问题排查与结构优化
一、核心问题排查方向
尺寸不匹配的根源几乎都来自forward方法中各模块的维度衔接错误,重点检查以下几点:
- Embedding输出维度与Transformer注意力层的输入维度是否一致
- 注意力层输出后,序列维度是否与后续全连接层的输入/目标张量维度对齐
- 训练时输入序列、目标序列的形状是否和模型输出匹配
二、典型错误代码与修正示例
1. 常见错误模型实现(尺寸不匹配根源)
import torch import torch.nn as nn class BadSingleHeadTransformer(nn.Module): def __init__(self, vocab_size, embed_dim, num_classes): super().__init__() self.embedding = nn.Embedding(vocab_size, embed_dim) self.attn = nn.MultiheadAttention(embed_dim, num_heads=1, batch_first=True) self.fc = nn.Linear(embed_dim, num_classes) def forward(self, x): x = self.embedding(x) # 输出形状: [batch_size, seq_len, embed_dim] attn_out, _ = self.attn(x, x, x) # 输出形状: [batch_size, seq_len, embed_dim] # 错误:直接将序列维度张量传给全连接,若目标是[batch_size, num_classes],维度完全不匹配 output = self.fc(attn_out) return output
2. 修正后的模型实现
import torch import torch.nn as nn class SingleHeadTransformer(nn.Module): def __init__(self, vocab_size, embed_dim, seq_len, num_classes, pad_idx=0): super().__init__() self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=pad_idx) # 必须加入位置编码,Transformer无法天然捕捉序列顺序 self.pos_encoding = nn.Parameter(torch.randn(1, seq_len, embed_dim)) self.attn = nn.MultiheadAttention(embed_dim, num_heads=1, batch_first=True) # 残差连接+LayerNorm,避免梯度消失,提升训练稳定性 self.norm = nn.LayerNorm(embed_dim) # 全连接层输入为聚合后的序列特征,维度匹配目标张量 self.fc = nn.Linear(embed_dim, num_classes) def forward(self, x, padding_mask=None): batch_size, seq_len = x.shape # 嵌入层+位置编码 x = self.embedding(x) + self.pos_encoding[:, :seq_len, :] # 计算注意力,传入padding mask避免关注无效pad token attn_out, _ = self.attn(x, x, x, key_padding_mask=padding_mask) # 残差连接+归一化 x = self.norm(x + attn_out) # 聚合序列维度:用均值池化(也可选用cls token、最大池化) x = torch.mean(x, dim=1) # 输出形状: [batch_size, embed_dim] # 全连接层输出匹配目标张量维度 output = self.fc(x) # 输出形状: [batch_size, num_classes] return output
三、训练循环的维度对齐检查
- 分类任务:输入张量
x形状为[batch_size, seq_len],目标张量y形状为[batch_size]或[batch_size, num_classes],需确保模型输出维度与之匹配 - 序列生成任务:目标张量形状为
[batch_size, target_seq_len],此时模型无需聚合序列维度,全连接层需输出[batch_size, target_seq_len, vocab_size],对应每个时间步的预测
四、针对性报错处理
若报错为RuntimeError: Expected target size [N, C], got [N, S]:
- 检查目标张量是否完成正确转换(比如分类任务中是否将序列标签压缩为单标签)
- 确认模型
forward中是否对Transformer输出的序列维度做了聚合操作(如均值池化、取cls token)
五、结构优化建议
- 强制加入位置编码:这是Transformer的核心组件,缺失会导致模型完全无法学习序列顺序特征
- 增加Feed Forward Network(FFN):标准Transformer Encoder包含Attention+FFN,FFN可进一步提取特征,结构为
nn.Linear(embed_dim, 4*embed_dim) -> nn.ReLU() -> nn.Linear(4*embed_dim, embed_dim) - 完善padding mask处理:避免注意力机制关注padding token,提升计算有效性
- 加入 dropout:在Embedding层、Attention输出后加入dropout,防止过拟合
内容的提问来源于stack exchange,提问作者armandp
相关产品推荐
相关产品推荐

