PyTorch实现Luong Attention遇阻:General模式无法正常学习求排查
我仔细对照Luong等人2015年论文里的general注意力模式,梳理了你的代码,发现几个可能导致模型无法有效学习的关键点:
1. 语法错误:缺失右括号
你的代码中这一行少了一个闭合的右括号:
out_hc = F.tanh(self.Whc(torch.cat([hidden[0], context], dim=1))
必须修正为:
out_hc = F.tanh(self.Whc(torch.cat([hidden[0], context], dim=1)))
虽然你提到代码可运行,但大概率是粘贴时遗漏了这个括号——如果实际运行时存在这个错误,会导致张量运算维度异常,间接破坏模型的梯度传播,影响收敛。
2. 批量兼容性不足:误用torch.mm导致单样本依赖
你的代码目前仅支持单样本(batch_size=1)训练,因为使用了torch.mm(仅处理2D张量)。单样本训练不仅收敛速度慢,而且稳定性极差,很容易出现无法学习的情况。
针对general注意力的批量计算,应该改用torch.bmm(批量矩阵乘法),调整注意力计算部分的代码如下:
# 调整decoder隐藏状态的形状:从(num_layers, batch_size, hidden_size)转为(batch_size, 1, hidden_size) hidden_attn = self.attn(hidden[0]).unsqueeze(1) # shape: (batch_size, 1, hidden_size) # 将encoder输出转置为适配批量矩阵乘法的形状:(batch_size, hidden_size, seq_len) encoder_outputs_trans = encoder_outputs.permute(1, 2, 0) # 假设原encoder_outputs是(seq_len, batch_size, hidden_size) # 批量计算注意力分数:(batch_size,1,hidden_size) @ (batch_size,hidden_size,seq_len) -> (batch_size,1,seq_len) attn_prod = torch.bmm(hidden_attn, encoder_outputs_trans) # 对seq_len维度做softmax,得到注意力权重 attn_weights = F.softmax(attn_prod, dim=2).squeeze(1) # shape: (batch_size, seq_len) # 批量计算上下文向量:(batch_size,1,seq_len) @ (batch_size,seq_len,hidden_size) -> (batch_size, hidden_size) context = torch.bmm(attn_weights.unsqueeze(1), encoder_outputs.permute(1,0,2)).squeeze(1)
修改后代码支持批量训练,能大幅提升模型的学习效率和稳定性。
3. 初始隐藏状态的关键错误
在Seq2Seq模型中,decoder的初始隐藏状态必须设置为encoder的最后隐藏状态(这是Luong论文里明确的设定)。如果你的训练代码中是随机初始化decoder的hidden,或者传入了错误的初始值,模型会完全丢失编码器的上下文信息,根本无法学习到有效的输入输出映射。
确保训练时传入decoder的初始hidden是encoder输出的hidden(比如单向GRU的话,取encoder_hidden[-1:]即可)。
4. 激活函数的版本兼容性
PyTorch中F.tanh已被标记为废弃,建议改用torch.tanh,避免后续版本的兼容性问题:
out_hc = torch.tanh(self.Whc(torch.cat([hidden[0], context], dim=1)))
额外验证建议
训练过程中可以打印或可视化注意力权重,观察是否随着训练逐步聚焦到相关的encoder输入位置。如果注意力权重始终均匀分布,说明注意力机制没有正常工作,大概率是上述某个问题导致的。
内容的提问来源于stack exchange,提问作者zyxue

