注意力机制入门:模型如何聚焦关键

人在读一句话时,并不会平均地对待每个字:理解「他去了银行取钱」里的「银行」,我们会下意识联系「取钱」而不是「河岸」。注意力机制(attention)把这种「按相关性聚焦」的能力搬进了神经网络:模型处理某个位置时,可以依据内容本身,动态决定该多看输入里的哪些部分、少看哪些部分。

Q、K、V 分别是什么

注意力围绕三个向量角色展开:

  • 查询(Query, Q):当前位置「我想找什么」的表示。
  • 键(Key, K):每个输入位置「我能提供什么」的标签,用来和查询比对相关性。
  • 值(Value, V):每个输入位置「真正携带的信息」,按相关性加权后汇总输出。

计算分两步:先用 Q 与各个 K 算相似度,得到一组权重;再用这些权重对 V 做加权求和。相似度高的位置,其 V 在结果里占的比重就大。

缩放点积注意力

最常用的是缩放点积注意力(scaled dot-product attention)。对一组查询、键、值,先算 Q 与 K 的点积衡量相关性,除以根号 d_k 做缩放以稳定梯度,经 softmax 得到权重,再乘 V:

import numpy as np

def scaled_dot_product_attention(Q, K, V):
    d_k = Q.shape[-1]
    scores = Q @ K.transpose(-1, -2) / np.sqrt(d_k)   # 相关性分数
    weights = np.exp(scores) / np.sum(np.exp(scores), axis=-1, keepdims=True)  # softmax
    return weights @ V                                 # 按权重汇总值

# 示例:3 个查询位置,每个看 4 个键/值,向量维度 5
Q = np.random.randn(3, 4, 5)
K = np.random.randn(3, 4, 5)
V = np.random.randn(3, 4, 5)
out = scaled_dot_product_attention(Q, K, V)
print(out.shape)   # (3, 4, 5)

这段代码演示了核心矩阵运算:一次矩阵乘法就能并行算出所有位置对之间的注意力权重,再与 V 相乘得到上下文相关的输出。

自注意力与交叉注意力

  • 自注意力(self-attention):Q、K、V 都来自同一序列。每个位置都能直接关注序列内其他位置,因此擅长捕捉长距离依赖,例如句首的词影响句尾的理解。
  • 交叉注意力(cross-attention):Q 来自一个序列,K、V 来自另一个序列。典型用于编码器到解码器的衔接,例如在翻译时,解码端用查询去检索编码端源语言里相关的词。

两者公式完全相同,区别只在 Q、K、V 的来源。

为什么注意力有用

注意力的价值体现在几个方面:

  • 长距离依赖:无论两个位置相隔多远,自注意力一步即可建立直接联系,而不必像循环网络那样逐层传递。
  • 可解释性:注意力权重可视化为「这个词在看哪些词」,为模型决策提供粗略的线索。
  • 并行高效:相关性的计算可完全矩阵化,充分利用 GPU,相比循环结构更易扩展。

当然,注意力并非万能:它对序列长度的复杂度是平方级,超长序列会带来显存与计算压力,这也是后续稀疏注意力、线性注意力等变体出现的动机。

小结

注意力机制通过 Query、Key、Value 三者的交互,让模型按相关性动态聚焦输入的不同部分。缩放点积注意力以「点积除根号维度再 softmax 加权 V」为核心,可完全矩阵化并行;自注意力处理同一序列、交叉注意力连接两个序列。它因擅长长距离依赖、可解释与易并行而被 Transformer 广泛采用,但平方级复杂度也催生了众多高效变体。

参考与延伸阅读

  • Vaswani et al. Attention Is All You Need(NeurIPS 2017)。提出 Transformer 与缩放点积注意力的原作。已核验。
  • 若想动手实验,Hugging Face 的 Transformers 文档与 The Annotated Transformer 教程给出了可运行的实现;具体接口以官方文档为准。待核实。
本文累计阅读