TensorFlow中三种Attention层的差异及MultiHeadAttention底层实现
TensorFlow中三种Attention层的区别、MultiHeadAttention手动实现及指定教程逻辑解析
一、tf.keras.layers.Attention、MultiHeadAttention、AdditiveAttention的核心区别
- tf.keras.layers.Attention:基础的缩放点积注意力实现,遵循Transformer原论文逻辑。计算时先对query和key做点积,再除以key维度的平方根避免梯度爆炸/消失,最后用softmax生成权重并与value加权求和。属于单头注意力,适合无需多头拆分的简单场景。
- tf.keras.layers.MultiHeadAttention:基于缩放点积注意力扩展,将query、key、value拆分为多个子空间(多头),每个头独立计算注意力,最后拼接所有头的输出并做线性变换得到结果。多头设计能让模型捕捉不同维度的语义信息,是Transformer架构的核心组件。
- tf.keras.layers.AdditiveAttention:又称Bahdanau注意力,属于加性注意力。不直接对query和key做点积,而是通过共享前馈网络将两者映射到同一维度后相加,经tanh激活和线性层得到注意力权重。更适合query与key维度不同的场景,也是机器翻译早期常用的注意力机制,你提到的教程即采用此实现。
二、用基础层手动实现MultiHeadAttention
核心思路:先通过Dense层投影query、key、value到目标维度,拆分多头后计算每个头的缩放点积注意力,拼接结果后做最终线性变换,可选加入残差连接和层归一化(Transformer标准结构)。
import tensorflow as tf class CustomMultiHeadAttention(tf.keras.layers.Layer): def __init__(self, d_model, num_heads): super().__init__() self.num_heads = num_heads self.d_model = d_model # 确保模型维度可被头数整除 assert d_model % num_heads == 0 self.depth = d_model // num_heads # 定义投影层:query、key、value各一个,最终输出一个 self.wq = tf.keras.layers.Dense(d_model) self.wk = tf.keras.layers.Dense(d_model) self.wv = tf.keras.layers.Dense(d_model) self.dense = tf.keras.layers.Dense(d_model) # 残差连接与层归一化组件 self.add = tf.keras.layers.Add() self.layernorm = tf.keras.layers.LayerNormalization(epsilon=1e-6) def split_heads(self, x, batch_size): # 维度转换:(batch_size, seq_len, d_model) -> (batch_size, num_heads, seq_len, depth) x = tf.reshape(x, (batch_size, -1, self.num_heads, self.depth)) return tf.transpose(x, perm=[0, 2, 1, 3]) def scaled_dot_product_attention(self, q, k, v, mask=None): # 计算缩放点积注意力 matmul_qk = tf.matmul(q, k, transpose_b=True) # 缩放操作,避免点积值过大导致softmax梯度消失 dk = tf.cast(tf.shape(k)[-1], tf.float32) scaled_attention_logits = matmul_qk / tf.math.sqrt(dk) # 应用mask(如padding mask、前瞻mask) if mask is not None: scaled_attention_logits += (mask * -1e9) # 生成注意力权重 attention_weights = tf.nn.softmax(scaled_attention_logits, axis=-1) # 加权求和得到注意力输出 output = tf.matmul(attention_weights, v) return output, attention_weights def call(self, inputs, mask=None): batch_size = tf.shape(inputs['query'])[0] # 第一步:投影query、key、value q = self.wq(inputs['query']) k = self.wk(inputs['key']) v = self.wv(inputs['value']) # 第二步:拆分多头 q = self.split_heads(q, batch_size) k = self.split_heads(k, batch_size) v = self.split_heads(v, batch_size) # 第三步:计算每个头的注意力 scaled_attention, attention_weights = self.scaled_dot_product_attention(q, k, v, mask) # 第四步:拼接多头输出 scaled_attention = tf.transpose(scaled_attention, perm=[0, 2, 1, 3]) concat_attention = tf.reshape(scaled_attention, (batch_size, -1, self.d_model)) # 第五步:最终线性变换 output = self.dense(concat_attention) # 残差连接+层归一化(可选,符合Transformer标准结构) output = self.add([inputs['query'], output]) output = self.layernorm(output) return output, attention_weights
三、指定教程中注意力层的操作逻辑
教程中的注意力层为AdditiveAttention(Bahdanau注意力),用于机器翻译解码器,核心是让解码器每一步都能关注编码器的全部输出,具体逻辑如下:
- 输入接收:解码器前一步隐藏状态、编码器所有输出、编码器padding mask(避免关注padding部分)。
- 维度适配:将解码器隐藏状态通过Dense层映射到与编码器输出相同的维度,扩展维度以匹配编码器的序列长度;编码器输出也通过共享Dense层做同样映射。
- 注意力分数计算:将映射后的解码器隐藏状态与编码器输出相加,经
tanh激活得到中间特征,再通过输出维度为1的Dense层压缩为每个时间步的注意力分数。 - 权重归一化:对注意力分数应用padding mask(将padding位置的分数设为负无穷),再用
softmax归一化,得到编码器每个时间步的注意力权重。 - 上下文向量生成:用注意力权重对编码器输出做加权求和,得到上下文向量,该向量会与解码器当前步输入拼接,作为解码器下一步的输入。
这套逻辑通过加性方式计算解码器与编码器的关联程度,让解码器生成每个词时都能聚焦到编码器中最相关的输入部分,解决了长序列翻译的依赖问题。
内容的提问来源于stack exchange,提问作者user15101374
相关产品推荐
相关产品推荐

