人工智能

Transformer 注意力机制详解:从 Seq2Seq 到 Self-Attention

"Attention Is All You Need" 之后,注意力成了大模型的通用底座。本文从 Seq2Seq 的信息瓶颈讲起,逐行拆解 QKV 计算与多头注意力,并给出一个最小可运行的 NumPy 实现。

在 Transformer 出现之前,机器翻译的主流是编码器-解码器加循环网络。它有一个难以回避的结构性问题: 编码器把整句话压成一个固定长度的向量,解码器只能看这一个向量。句子越长,信息丢得越多。

注意力的本质是一次加权求和:每个输出位置,都可以自主决定”该看输入的哪些位置、看多重”。

1. Seq2Seq 的信息瓶颈

设输入长度为 T,隐藏状态维度为 d。RNN 编码器要在读完最后一个词之后, 把整句话的信息全部塞进一个 d 维向量里。当 T 远大于 d 时, 这几乎是不可能完成的任务——长句翻译质量断崖式下跌就是这么来的。

更糟的是,RNN 必须串行计算:处理第 t 个词时依赖第 t−1 个隐藏状态, 无法并行。训练一个长序列模型,GPU 有一大半时间在等前一步算完。

2. 注意力:让解码器”回头看”

注意力机制的思路很直接:别压成向量了,把编码器每个位置的隐藏状态都留着, 解码每生成一个词,就去这些状态里按相关度加权取一次。

αt,i = softmax(score(ht, h̄i)),  ct = Σi αt,i · h̄i (1)

权重 α 由相似度分数归一化得到,c 是加权后的上下文向量。 这样一来,信息通道从固定长度变成了随输入线性增长,长句不再是瓶颈。

3. 自注意力与 QKV

Transformer 更进一步:既然注意力这么好用,为什么不用在输入序列自己身上? 于是有了自注意力——序列中每个位置都去关注同一序列的所有位置。

为了让”查询”和”被查询”的角色分离,同一份输入被投影成三组向量:Query(我在找什么)、 Key(我是什么)、Value(我能提供什么):

Attention(Q, K, V) = softmax( Q·KT / √dk ) · V (2)

关键在于除以 √dk。点积的方差随维度增大而增大,不缩放的话 softmax 会饱和成近似 one-hot, 梯度几乎消失,训练直接卡住。这个看似不起眼的因子,是 Transformer 能堆深的必要条件之一。

4. 多头注意力与位置编码

单组 QKV 只能捕捉一种”关注模式”,多头则把 d 维切分成 h 组并行计算, 每组独立学习不同的关系:有的头盯语法结构,有的头盯指代关系。最后拼接再线性变换回去。

组件作用
多头注意力在不同子空间并行捕捉多种依赖关系
位置编码补回序列顺序信息,因为自注意力本身与位置无关
残差连接让梯度直通,支撑深层堆叠
LayerNorm稳定各层输出分布,加速收敛

位置编码是必须的:自注意力对输入做的是集合运算,打乱词序结果完全一样。 原始论文用的是不同频率的正余弦函数,现在更多模型改用可学习的位置向量或相对位置编码。

5. 最小可运行实现

下面这段代码只依赖 NumPy,把缩放点积注意力与多头拆分完整实现了一遍,可以直接跑通并观察注意力矩阵的形状。

import numpy as np

def softmax(x, axis=-1):
    x = x - x.max(axis=axis, keepdims=True)   # 数值稳定
    e = np.exp(x)
    return e / e.sum(axis=axis, keepdims=True)

def scaled_dot_attention(Q, K, V, mask=None):
    dk = Q.shape[-1]
    scores = Q @ K.transpose(0, 2, 1) / np.sqrt(dk)
    if mask is not None:
        scores = np.where(mask, scores, -1e9)   # 屏蔽未来位置
    weights = softmax(scores)
    return weights @ V, weights
class MultiHeadAttention:
    def __init__(self, d_model, n_heads, rng):
        assert d_model % n_heads == 0
        self.h = n_heads
        self.dk = d_model // n_heads
        scale = 1.0 / np.sqrt(d_model)
        # 三个投影矩阵:把输入分别变成 Q、K、V
        self.Wq = rng.normal(0, scale, (d_model, d_model))
        self.Wk = rng.normal(0, scale, (d_model, d_model))
        self.Wv = rng.normal(0, scale, (d_model, d_model))
        self.Wo = rng.normal(0, scale, (d_model, d_model))

    def split_heads(self, x):
        B, T, D = x.shape
        # (B,T,D) -> (B,h,T,dk),让每个头看到完整序列
        return x.reshape(B, T, self.h, self.dk).transpose(0, 2, 1, 3)

    def forward(self, x):
        Q = self.split_heads(x @ self.Wq)
        K = self.split_heads(x @ self.Wk)
        V = self.split_heads(x @ self.Wv)
        out, attn = scaled_dot_attention(Q, K, V)

        # 拼回 (B,T,D),再过输出投影
        B, h, T, dk = out.shape
        out = out.transpose(0, 2, 1, 3).reshape(B, T, h * dk)
        return out @ self.Wo, attn

跑起来看形状:输入 (2, 10, 64)(2 条样本、长度 10、维度 64),4 个头, 输出仍是 (2, 10, 64),而注意力矩阵是 (2, 4, 10, 10)—— 每个头、每条样本都有一张自己的注意力分布图。

rng = np.random.default_rng(0)
mha = MultiHeadAttention(d_model=64, n_heads=4, rng=rng)
x = rng.normal(0, 1, (2, 10, 64))

out, attn = mha.forward(x)
print(out.shape)    # (2, 10, 64)
print(attn.shape)   # (2, 4, 10, 10)
print(attn[0, 0, 0].sum())   # 1.0,每行权重和为 1

6. 总结

从 Seq2Seq 到 Transformer,注意力完成了从”辅助机制”到”唯一主角”的转变。 核心公式只有一条:缩放点积注意力;多头负责并行捕捉多种关系,位置编码补回顺序,残差与 LayerNorm 保证能堆深。

理解这条主干后,无论是 BERT 的双向编码、GPT 的因果掩码,还是各种高效注意力变体, 差别都只在掩码方式、连接结构和计算顺序上。真正难的从来不是公式,而是把它在工程上跑稳。