TensorFlow MultiHeadAttention掩码销毁警告原因及解决咨询
TransformerBlock中EinsumDense掩码警告的原因与解决方法
问题背景
我正在使用tensorflow==2.16.1构建Transformer模型,自定义TransformerBlock实现如下:
# Import TensorFlow and Keras for building and training neural network models import tensorflow as tf from tensorflow.keras.layers import ( Dense, LayerNormalization, MultiHeadAttention, Dropout, ) class TransformerBlock(tf.keras.layers.Layer): def __init__(self, embed_dim, num_heads, ff_dim, rate=0.1, **kwargs): super(TransformerBlock, self).__init__(**kwargs) self.embed_dim = embed_dim self.num_heads = num_heads self.ff_dim = ff_dim self.rate = rate self.att = None self.ffn = None self.layernorm1 = None self.layernorm2 = None self.dropout1 = None self.dropout2 = None def build(self, input_shape): self.att = MultiHeadAttention(num_heads=self.num_heads, key_dim=self.embed_dim) self.ffn = tf.keras.Sequential( [Dense(self.ff_dim, activation="relu"), Dense(self.embed_dim)] ) self.layernorm1 = LayerNormalization(epsilon=1e-6) self.layernorm2 = LayerNormalization(epsilon=1e-6) self.dropout1 = Dropout(self.rate) self.dropout2 = Dropout(self.rate) super(TransformerBlock, self).build(input_shape) def call(self, inputs, training, padding_mask=None, causal_mask=True, qa=False): mask = None seq_len = tf.shape(inputs)[1] batch_size = tf.shape(inputs)[0] if padding_mask is not None: padding_mask_reshaped = tf.cast( tf.reshape(padding_mask, (batch_size, 1, seq_len)), dtype=tf.float32 ) mask = tf.broadcast_to( padding_mask_reshaped, (batch_size, seq_len, seq_len) ) attn_output = self.att( inputs, inputs, attention_mask=mask, use_causal_mask=True ) attn_output = self.dropout1(attn_output, training=training) out1 = self.layernorm1(inputs + attn_output) ffn_output = self.ffn(out1) ffn_output = self.dropout2(ffn_output, training=training) out2 = self.layernorm2(out1 + ffn_output) return out2
每次运行该TransformerBlock时,都会收到如下警告:
lib/python3.11/site-packages/keras/src/layers/layer.py:877: UserWarning: Layer 'value' (of type EinsumDense) was passed an input with a mask attached to it. However, this layer does not support masking and will therefore destroy the mask information. Downstream layers will not see the mask.
实际测试中,传入padding_mask并设置use_causal_mask=True时模型性能符合预期(比如use_causal_mask=False时性能异常优异,说明causal掩码确实生效);将causal掩码与padding_mask合并后传入attention_mask参数也会出现相同性能变化。
警告产生的原因
MultiHeadAttention内部会用EinsumDense层处理Q/K/V的投影转换,而你输入的inputs附带了Keras内置掩码信息(大概率来自上游Embedding层或输入预处理时自动生成的掩码)。EinsumDense层本身不支持Keras的掩码传递机制,当它接收到带掩码的输入时,会直接丢弃该掩码信息,因此触发这个警告,提示下游层无法获取原有的自动掩码。- 注意:这个警告不影响你手动传入的
attention_mask或use_causal_mask的生效——这也是为什么你能看到掩码对性能的实际影响,警告只是针对Keras内置的掩码传递逻辑,而非你手动设置的注意力掩码。
消除警告的方法
方法1:手动剥离输入的掩码信息
在TransformerBlock的call方法开头,通过Lambda层剥离输入的掩码,避免传递给MultiHeadAttention内部的EinsumDense层:
def call(self, inputs, training, padding_mask=None, causal_mask=True, qa=False): # 剥离输入的掩码信息 inputs = tf.keras.layers.Lambda(lambda x: x)(inputs) mask = None # 后续原有逻辑保持不变 seq_len = tf.shape(inputs)[1] batch_size = tf.shape(inputs)[0] # ... 剩余代码省略
方法2:调用MultiHeadAttention时显式关闭输入掩码传递
如果你不需要依赖Keras的自动掩码传递,可以在调用self.att时,显式设置mask=None(这里的mask指输入的Keras内置掩码,不是你手动构建的注意力掩码):
attn_output = self.att( inputs, inputs, attention_mask=mask, use_causal_mask=True, mask=None # 显式关闭输入的掩码传递 )
方法3:直接屏蔽该类警告
如果确认手动掩码逻辑完全正常,只是不想看到警告,可以通过Python的警告过滤机制屏蔽:
import warnings from tensorflow.keras.src.layers.layer import UserWarning warnings.filterwarnings("ignore", category=UserWarning, message="Layer 'value' \(of type EinsumDense\) was passed an input with a mask attached to it.*")
内容的提问来源于stack exchange,提问作者Stod
相关产品推荐
相关产品推荐

