You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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注意力),用于机器翻译解码器,核心是让解码器每一步都能关注编码器的全部输出,具体逻辑如下:

  1. 输入接收:解码器前一步隐藏状态、编码器所有输出、编码器padding mask(避免关注padding部分)。
  2. 维度适配:将解码器隐藏状态通过Dense层映射到与编码器输出相同的维度,扩展维度以匹配编码器的序列长度;编码器输出也通过共享Dense层做同样映射。
  3. 注意力分数计算:将映射后的解码器隐藏状态与编码器输出相加,经tanh激活得到中间特征,再通过输出维度为1的Dense层压缩为每个时间步的注意力分数。
  4. 权重归一化:对注意力分数应用padding mask(将padding位置的分数设为负无穷),再用softmax归一化,得到编码器每个时间步的注意力权重。
  5. 上下文向量生成:用注意力权重对编码器输出做加权求和,得到上下文向量,该向量会与解码器当前步输入拼接,作为解码器下一步的输入。

这套逻辑通过加性方式计算解码器与编码器的关联程度,让解码器生成每个词时都能聚焦到编码器中最相关的输入部分,解决了长序列翻译的依赖问题。

内容的提问来源于stack exchange,提问作者user15101374

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.30 13:02:34