TensorFlow Keras中Transformer不同批次可变长度输入的掩码构建问题
解决Keras MultiHeadAttention中attention_mask的形状匹配问题
核心问题梳理
你当前的掩码设计混淆了序列长度维度和特征维度,导致形状不匹配。你的输入数据形状是[B, T, C](B=批量大小,T=序列长度=4,C=特征维度=3),而掩码应该针对无效的序列行(比如含nan的行),而非特征维度。
正确的掩码构建步骤
1. 生成基础序列掩码
首先从输入数据中识别无效序列行,生成形状为[B, T]的基础掩码(1表示有效行,0表示无效行):
import numpy as np import tensorflow as tf # 假设你的输入数据为data,形状[10,4,3] mask = ~np.isnan(data).any(axis=-1) # 一行中只要有nan就标记为无效(False) mask = mask.astype(np.int32) # 转换为int类型,最终形状[10,4]
2. 转换为MultiHeadAttention兼容的形状
自注意力场景下(Query=Key=Value),需要将基础掩码扩展为[B, T, T]或[B, num_heads, T, T],确保每个Query位置只能关注到有效的Key位置:
- 方法一:生成三维掩码
[B, T, T]
# 确保Query和Key位置都有效时才允许注意力计算 attention_mask = tf.expand_dims(mask, 1) & tf.expand_dims(mask, 2) # 形状[B,T,T]
- 方法二:生成四维掩码适配多头注意力(自动广播到所有head)
# 先得到[B,T,T],再扩展head维度 attn_mask_3d = tf.matmul(tf.expand_dims(mask, -1), tf.expand_dims(mask, 1)) attention_mask = tf.expand_dims(attn_mask_3d, 1) # 形状[B,1,T,T],适配num_heads维度
模型适配修正
不能将掩码作为固定参数传入模型,需将其作为输入之一,与每个batch的样本一一对应:
def transformer_encoder(inputs, head_size, num_heads, ff_dim, dropout=0.0, mask=None): x = layers.LayerNormalization(epsilon=1e-6)(inputs) x = layers.MultiHeadAttention( key_dim=head_size, num_heads=num_heads, dropout=dropout )(x, x, attention_mask=mask) x = layers.Dropout(dropout)(x) res = x + inputs x = layers.LayerNormalization(epsilon=1e-6)(res) x = layers.Conv1D(filters=ff_dim, kernel_size=1, activation="relu")(x) x = layers.Dropout(dropout)(x) x = layers.Conv1D(filters=inputs.shape[-1], kernel_size=1)(x) return x + res def build_model( n_classes, input_shape, head_size, num_heads, ff_dim, num_transformer_blocks, mlp_units, dropout=0.0, mlp_dropout=0.0, ) -> keras.Model: inputs = keras.Input(shape=input_shape) mask_input = keras.Input(shape=(input_shape[0],)) # 掩码输入,形状[T] x = inputs for _ in range(num_transformer_blocks): # 实时转换掩码为兼容形状 attn_mask_3d = tf.matmul(tf.expand_dims(mask_input, -1), tf.expand_dims(mask_input, 1)) attn_mask = tf.expand_dims(attn_mask_3d, 1) x = transformer_encoder(x, head_size, num_heads, ff_dim, dropout, mask=attn_mask) x = layers.GlobalAveragePooling2D(data_format="channels_first")(x) for dim in mlp_units: x = layers.Dense(dim, activation="relu")(x) x = layers.Dropout(mlp_dropout)(x) outputs = layers.Dense(n_classes, activation="softmax")(x) return keras.Model(inputs=[inputs, mask_input], outputs=outputs)
训练时的使用方式
将输入数据和对应掩码一起传入模型:
# 假设labels是你的标签数据 model.fit([data, mask], labels, batch_size=5, epochs=10)
关键注意点
- 掩码始终针对序列长度维度,而非特征维度,不要生成
[B,T,C]形状的掩码 - 自注意力场景下,掩码需要覆盖Query和Key的所有位置组合,因此形状需扩展为
[B,T,T]或适配多头的四维形状 - 掩码必须与输入数据一一对应,每个batch的掩码随样本动态传入,不能预先固定
内容的提问来源于stack exchange,提问作者Jorge Morgado
相关产品推荐
相关产品推荐

