从零实现Transformer遇形状错误:编码器解码器输入维度不匹配
解决Transformer编码器-解码器形状不匹配的RuntimeError
这个RuntimeError的核心是:代码尝试将总元素数为2048的张量重塑为[1, 40, 64](总元素数2560),两者维度不匹配,问题出在编码器输出传递到解码器的环节。以下是具体排查和解决方法:
1. 修正硬编码的序列长度
- 问题根源:你可能在解码器中硬编码了序列长度为40,但编码器实际处理的序列长度是32(1×32×64=2048,正好对应错误中的输入大小)。
- 解决:不要固定写死序列长度,改用动态获取方式。比如在Transformer主类的forward函数中,从编码器输出中提取实际序列长度:
# 获取编码器输出的实际序列长度 src_seq_len = encoder_output.size(1) # 生成对应形状的掩码(如果需要) src_mask = torch.zeros(1, src_seq_len, src_seq_len).to(encoder_output.device) # 传递给解码器时使用动态获取的长度 decoder_out = self.decoder(tgt, encoder_output, src_mask)
2. 统一模型的特征维度配置
- 问题根源:编码器和解码器的
d_model(特征维度)参数不一致,或者中间层(如嵌入层、前馈层)私自修改了维度,导致输出总元素数偏差。 - 解决:确保编码器、解码器、嵌入层、注意力层的
d_model参数完全相同,比如都设置为64,全程不要随意改变特征维度。
3. 检查编码器输出的形状
- 问题根源:编码器的forward函数可能被错误添加了降维操作(如
flatten()、torch.mean(dim=1)),把原本[batch_size, seq_len, d_model]的三维张量改成了一维或二维。 - 解决:在编码器forward函数末尾打印输出形状,确认是
[1, 32, 64]这类三维格式,移除所有不必要的降维操作。
4. 验证解码器的上下文输入形状
- 问题根源:解码器的交叉注意力层期望的上下文张量形状是
[batch_size, src_seq_len, d_model],但你传递的张量维度不符。 - 解决:在Transformer主类中,分别打印编码器输出和解码器接收前的张量形状,确保两者的batch_size、序列长度、特征维度完全匹配。
内容的提问来源于stack exchange,提问作者user20372902
相关产品推荐
相关产品推荐

