返回技术文章

注意力机制到底在算什么

  • 深度学习
  • Transformer
  • 注意力

示例文章,用于演示分页与列表排版,请替换成你自己的内容。

注意力机制的公式很短,但每一项为什么在那里,值得拆开看一遍。

从一次加权平均说起

给定查询 qq 和一组键值对 {(ki,vi)}\{(k_i, v_i)\},注意力的输出是值的加权平均:

Attention(q,K,V)=iαivi,αi=exp(qki)jexp(qkj)\mathrm{Attention}(q, K, V) = \sum_{i} \alpha_i v_i, \qquad \alpha_i = \frac{\exp(q \cdot k_i)}{\sum_j \exp(q \cdot k_j)}

权重 αi\alpha_i 由查询与键的内积决定,再经过 softmax 归一化。内积大表示「这个键和我相关」,于是对应的值被分到更大的权重。

为什么要除以 dk\sqrt{d_k}

缩放点积注意力里那个 dk\sqrt{d_k} 不是随手加的:

softmax(QKdk)V\mathrm{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}\right)V

当维度 dkd_k 变大时,内积的方差随维度线性增长,进入 softmax 前的数值会拉得很开,梯度会被推到接近饱和的区域。除以 dk\sqrt{d_k} 把方差拉回常数,训练才稳定。

多头是在并行看不同的关系

单个注意力头只能输出一种加权方式。多头把表示切成若干份,各自学一套 Q,K,VQ, K, V

MultiHead(Q,K,V)=Concat(head1,,headh)WO\mathrm{MultiHead}(Q, K, V) = \mathrm{Concat}(\mathrm{head}_1, \dots, \mathrm{head}_h) W^O

不同的头容易分化出不同的关注模式——有的盯相邻位置,有的盯句法上的依赖。这是经验观察,不是设计时的硬性保证。

复杂度

自注意力的计算量是 O(n2d)O(n^2 d)nn 是序列长度。序列翻倍,计算量翻四倍,这也是长上下文昂贵的原因。各种稀疏化、线性注意力的工作,本质上都在改写这个平方项。