TensorFlow 2.11.0中MultiHeadAttention指定attention_axes=2时的掩码维度错误
多注意力头掩码维度不兼容问题分析与解决
问题背景
期望对输入形状(4,5,20,64)中的第2个维度(长度为20的维度)应用自注意力,使用带mask_zero=True的Embedding层时,抛出维度不兼容错误:
{{function_node _wrapped__AddV2_device/job:localhost/replica:0/task:0/device:CPU:0}} Incompatible shapes: [4,5,2,20,20] vs. [4,5,1,5,20] [Op:AddV2]
Call arguments received by layer 'softmax_2' (type Softmax):
• inputs=tf.Tensor(shape=(4, 5, 2, 20, 20), dtype=float32)
• mask=tf.Tensor(shape=(4, 5, 1, 5, 20), dtype=bool)
当mask_zero设为False时代码可正常运行。
问题原因
- Embedding层的自动掩码:当
mask_zero=True时,Embedding层会生成形状为(4,5,20)的布尔掩码(对应输入的(batch, 5, 20)维度),并自动传递给后续的MultiHeadAttention层。 - 注意力维度的掩码适配错误:指定
attention_axes=2后,MultiHeadAttention会对(4,5,20,64)中的第2个轴(长度20的维度)计算自注意力,生成的注意力权重形状为(4,5,2,20,20)(batch, 5, heads, seq_len_q, seq_len_k)。但自动传递的掩码被默认处理为适配所有空间维度的格式,被扩展成了(4,5,1,5,20),和注意力权重的维度无法匹配,导致相加操作失败。
解决方法
方法1:手动调整掩码形状,匹配注意力维度
在call方法中手动获取Embedding层的掩码,通过维度扩展将其调整为能广播到注意力权重形状的格式:
import numpy as np import tensorflow as tf from keras import layers as tfl class Encoder(tfl.Layer): def __init__(self,): super().__init__() self.embed_layer = tfl.Embedding(4500, 64, mask_zero=True) self.attn_layer = tfl.MultiHeadAttention(num_heads=2, attention_axes=2, key_dim=16) def call(self, x): # Input shape: (4, 5, 20) x, embed_mask = self.embed_layer(x, return_mask=True) # 获取掩码,形状(4,5,20) # 扩展掩码维度:从(4,5,20)变为(4,5,1,1,20),适配注意力权重的(4,5,2,20,20)形状 attn_mask = tf.expand_dims(tf.expand_dims(embed_mask, axis=2), axis=2) x = self.attn_layer(query=x, key=x, value=x, attention_mask=attn_mask) return x eg_input = tf.constant(np.random.randint(0, 150, (4, 5, 20))) enc = Encoder() enc(eg_input)
方法2:禁用自动掩码传递(若无需掩码)
如果不需要基于0值的掩码,可以直接将Embedding层的mask_zero设为False,避免自动生成和传递掩码,代码即可正常运行。
内容的提问来源于stack exchange,提问作者Vidyadhar Mudium
相关产品推荐
相关产品推荐

