Transformer解码器多头注意力掩码形状不匹配问题求助
Transformer解码器掩码匹配问题解决方案
问题核心
搭建Transformer解码器时,多头注意力的缩放点积输出形状为(64,8,64,64),但当前生成的掩码形状无法与之匹配,导致mul += -1e9 * tf.cast(mask,tf.float32)行出现形状不兼容错误。
掩码原理
解码器自注意力需要因果掩码(上三角结构),确保每个位置仅能访问序列中自身及之前的位置。缩放点积后的注意力分数矩阵形状为(batch_size, heads, seq_len_q, seq_len_k),掩码必须与该形状对齐,或通过广播机制适配。
当前代码的问题
- 原
create_mask生成的掩码形状为(batch_size, seq_len, seq_len),缺少heads维度,无法直接与注意力分数矩阵相加。 - 错误地使用
reshape_tensor处理掩码:掩码是序列级约束,所有注意力头共享同一掩码,无需按头拆分。
修正步骤
1. 重新实现掩码生成函数
生成适配多头注意力的掩码,利用广播机制兼容heads维度:
def create_mask(batch_size, seq_len): # 生成上三角掩码,1代表需要屏蔽的位置 mask = np.triu(np.ones((seq_len, seq_len)), 1) # 扩展维度为(1,1,seq_len,seq_len),方便批量广播 mask = np.expand_dims(np.expand_dims(mask, 0), 0) # 复制到对应批量大小,最终形状为(batch_size,1,seq_len,seq_len) mask = tf.tile(mask, [batch_size, 1, 1, 1]) return mask
该掩码可通过TensorFlow广播机制自动扩展heads维度,与(64,8,64,64)的注意力分数矩阵匹配。
2. 修改多头注意力层的掩码处理逻辑
移除对掩码的不必要重塑操作,直接传入缩放点积注意力层:
class MultiHeadAttentionLayer(Layer): # __init__部分保持不变 def call(self, query, key, val, return_score=False, mask=None): q_reshaped = self.reshape_tensor(self.Wq(query), self.heads, True) k_reshaped = self.reshape_tensor(self.Wk(key), self.heads, True) v_reshaped = self.reshape_tensor(self.Wv(val), self.heads, True) # 移除对mask的reshape_tensor处理 o_reshaped = self.attention(q_reshaped, k_reshaped, v_reshaped, self.dim_key, mask) output = self.reshape_tensor(o_reshaped, self.heads, False) if return_score: return self.Wo(output), K.sum(output, axis=1) return self.Wo(output)
3. 验证效果
修正后:
- 注意力分数矩阵形状:
(64,8,64,64) - 掩码形状:
(64,1,64,64)
通过广播机制,掩码会自动适配heads维度,成功与注意力分数矩阵相加,解决形状不匹配问题。
内容的提问来源于stack exchange,提问作者Shivam Singh
相关产品推荐
相关产品推荐

