跳到正文

Self-Attention 怎么计算 QKV?

注意力权重生成过程详解,从输入到加权求和

原题:请详细解释Self-Attention机制的工作原理,包括Query、Key、Value的计算过程以及注意力权重的生成方式。

模型架构 · 海尔真题

回答与解析

计算过程

给定输入X∈R^(L×d_model)

Q=XW_Q, K=XW_K, V=XW_V

Attention(Q,K,V)=softmax(QK^T/sqrt(d_k)+M)V

M可以表示causal mask、padding mask或其他可见性约束。除以sqrt(d_k)用于控制点积方差,避免维度增大时softmax过早饱和。每一行softmax得到当前query对各key的权重,再对V加权求和。

多头attention把Q、K、V投影到多个子空间,各头独立计算后拼接,并经线性输出投影恢复模型维度。注意力权重反映该前向计算中的加权系数,但不能直接当作因果解释或特征重要性的充分证据。

标准全局self-attention计算和显式分数存储随序列长度呈二次增长。FlashAttention减少中间矩阵的高带宽内存读写,但不改变全局attention的数学结果。KV Cache只在自回归解码中避免为历史token重复计算K/V;训练和prefill仍需处理整段输入,不能靠KV Cache消除其二次attention成本。

口语版讲法(约4分钟)

  • 从输入投影得到QKV
  • 解释缩放点积与mask
  • 说明softmax加权和多头
  • 限定注意力可解释性
  • 区分训练、prefill与decode复杂度

Self-Attention可以从一组token表示开始解释。输入矩阵X的形状是序列长度L乘模型维度d model,分别乘三组可学习矩阵得到Q、K、V。对某个位置而言,Q表示它要查询什么,K表示各位置可被匹配的特征,V是最终被加权汇总的内容。这些说法只是帮助理解,真正计算仍是线性投影和矩阵乘。

接着计算Q乘K转置,得到L乘L的相关分数。分数要除以根号d k,其中d k是每个头的key维度。若不缩放,维度增大时点积方差容易变大,softmax可能过早进入饱和区。然后加入mask。因果mask阻止当前位置看到未来token,padding mask排除补齐位置,二者的布尔语义和广播形状必须按具体框架确认。

对每个query所在的行做softmax后,得到对全部key的归一化权重,再与V相乘形成该位置的新表示。多头attention会把Q、K、V投影到多个较小子空间,各头独立完成这套计算,随后把结果拼接并通过输出线性层混合回模型维度。不同头可能学习不同关系,但不能预设每个头必然对应语法、实体或位置。

注意力权重也不能直接当成因果解释。它只能说明当前网络、当前输入和当前层里V被怎样加权;后续残差、MLP与其他层还会改变输出。两个模型可能有不同权重却给出相似预测,也可能改变某些权重但输出不按直觉变化。因此做可解释性分析时,还要结合干预、梯度或反事实实验。

复杂度方面,标准全局attention要形成所有token两两关系,主要计算随L平方增长,朴素实现还会显式存储L乘L分数。FlashAttention通过分块和在线softmax减少高带宽内存读写,数学上仍是精确全局attention,并没有把连接模式变稀疏。

KV Cache的边界也要说清。自回归decode每次只新增一个token,历史token的K和V不变,所以可以缓存,避免每步重新投影历史内容。但训练和prefill要同时处理整段输入,仍需计算整段QK关系。KV Cache降低的是解码中的重复计算,不会消除训练或prefill的二次attention。

若扩展到cross-attention,Q来自目标序列,K和V来自另一段memory,分数形状变为目标长度乘源长度。它和self-attention公式相似,但信息来源、mask和缓存策略不同。面试中先声明当前讲的是self-attention,能避免把两类连接混在一起。

全mask行需要特别处理,因为空集合上的softmax没有正常概率含义。生产实现应使用框架已验证的mask路径,并在因果、padding、混合精度和全mask样本上检查输出与梯度。

工程上我会先用小张量手算一行 attention,再和框架输出、mask 后概率和梯度对齐。全局 attention 还是稀疏 attention,要按上下文长度、延迟与显存做取舍;如果全 mask、混合精度或长序列出现 NaN,就应保留数值校验、告警与回退路径,不能只看最终 loss。

关键一句:为什么注意力权重不能直接作为模型决策的因果解释。

核验来源

  1. Attention Is All You Need
  2. Attention is not Explanation
  3. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness

面试官还可能这样问

  1. 问法 1 · 场景切入

    假设你在做一个电商搜索的排序模型,用户搜'苹果手机',你希望模型能自动关注商品标题里的'苹果'和'手机'这两个词,同时忽略'原装正品'这种修饰。这种token之间的关联度,底层是怎么算出来的?

  2. 问法 2 · 层层追问

    Transformer里每个token是怎么跟其他token交互的?……那它怎么知道该关注谁?这个关注的程度具体怎么量化?……对,从输入向量到得到权重分布,中间经过了哪些变换?

  3. 问法 3 · 直球架构

    讲一下Self-Attention中Query、Key、Value是怎么来的,以及注意力权重矩阵的计算过程,包括那个缩放因子为什么必要?

同模块相关题目