You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Transformer解码器多头注意力掩码形状不匹配问题求助

Transformer解码器掩码匹配问题解决方案

问题核心

搭建Transformer解码器时,多头注意力的缩放点积输出形状为(64,8,64,64),但当前生成的掩码形状无法与之匹配,导致mul += -1e9 * tf.cast(mask,tf.float32)行出现形状不兼容错误。

掩码原理

解码器自注意力需要因果掩码(上三角结构),确保每个位置仅能访问序列中自身及之前的位置。缩放点积后的注意力分数矩阵形状为(batch_size, heads, seq_len_q, seq_len_k),掩码必须与该形状对齐,或通过广播机制适配。

当前代码的问题

  1. 原create_mask生成的掩码形状为(batch_size, seq_len, seq_len),缺少heads维度,无法直接与注意力分数矩阵相加。
  2. 错误地使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.03 13:50:42