如何基于Transformer多注意力头实现双数值序列的交叉注意力?
用TensorFlow实现基于Multi-head Attention的交叉注意力
交叉注意力的核心逻辑是让一个输入序列(作为Query)去关注另一个输入序列(作为Key和Value)的信息,在TensorFlow的MultiHeadAttention层中,直接指定不同的输入作为query、key、value即可实现。针对你的两个(128,32)维度的输入,下面是完善的实现示例:
初始代码问题说明
Sequential模型仅适用于单输入单输出的线性结构,多输入场景需要使用Keras Functional API- 缺少Transformer结构中必备的残差连接和层归一化,这会影响模型的训练稳定性和性能
- 未正确完成模型的构建与输出定义
完整实现代码
import tensorflow as tf from tensorflow.keras import layers, Model # 定义输入层:两个输入都是(128, 32),即序列长度128,特征维度32 input1 = layers.Input(shape=(128, 32), name="input_sequence_1") input2 = layers.Input(shape=(128, 32), name="input_sequence_2") # 交叉注意力模块:用input1作为Query,input2作为Key和Value # 注意:key_dim需要能被num_heads整除,这里设置key_dim=16,num_heads=2,每个头处理16维特征 mha = layers.MultiHeadAttention(num_heads=2, key_dim=16, name="cross_attention") # 计算交叉注意力输出:query=input1, key=input2, value=input2 attention_output = mha(query=input1, value=input2, key=input2) # 添加残差连接 + 层归一化(Transformer标准组件) x = layers.Add()([input1, attention_output]) x = layers.LayerNormalization(epsilon=1e-6)(x) # 可选:添加后续的全连接或其他层(根据你的任务需求调整) x = layers.Dense(64, activation="relu")(x) x = layers.Dense(32)(x) # 构建完整模型 model = Model(inputs=[input1, input2], outputs=x) # 查看模型结构 model.summary() # 编译模型(根据任务类型选择损失函数和优化器) model.compile(optimizer="adam", loss="mse") # 测试模型输入输出 test_input1 = tf.random.normal((32, 128, 32)) # batch_size=32 test_input2 = tf.random.normal((32, 128, 32)) test_output = model([test_input1, test_input2]) print(f"输出维度:{test_output.shape}") # 应该是(32, 128, 32)
关键细节说明
- 交叉注意力的参数设置:
query是需要获取信息的序列,key和value是提供信息的序列,这里我们让input1关注input2的所有信息 - 维度匹配:
key_dim乘以num_heads最好等于输入的特征维度(这里16*2=32,刚好匹配输入的32维特征),这样每个注意力头可以均匀分配特征维度 - 残差与层归一化:这是Transformer结构的核心组件,残差连接保证信息不丢失,层归一化稳定训练过程,避免梯度消失或爆炸
- 任务适配:如果你的任务是分类、回归或其他任务,可以在模型末尾添加对应的输出层(比如分类任务加
Dense(num_classes, activation="softmax"))
内容的提问来源于stack exchange,提问作者Saurav Pahuja
相关产品推荐
相关产品推荐

