Skip to content

注意力机制

注意力机制让模型在处理一个元素时,能够根据相关性动态聚合其他元素的信息。它的核心是“查询 Query 去匹配键 Key,再按权重汇总值 Value”。

概念详解

给定查询矩阵 Q、键矩阵 K、值矩阵 V,缩放点积注意力为:

Attention(Q,K,V)=softmax(QKTdk)V

其中 dk 是 key 向量维度。除以 dk 是为了避免点积随维度增大而过大,导致 softmax 过于尖锐。

数学推导

对第 i 个 query,注意力分数为:

sij=qiTkj

缩放后:

s~ij=sijdk

权重为:

αij=exp(s~ij)mexp(s~im)

输出向量为:

oi=jαijvj

所以注意力本质上是对 value 的加权求和,权重由 query 和 key 的相似度决定。

应用代码

python
import torch
import torch.nn.functional as F

def scaled_dot_product_attention(Q, K, V, mask=None):
    d_k = Q.size(-1)
    scores = Q @ K.transpose(-2, -1) / (d_k ** 0.5)
    if mask is not None:
        scores = scores.masked_fill(mask == 0, float("-inf"))
    weights = F.softmax(scores, dim=-1)
    return weights @ V, weights

Q = torch.randn(2, 4, 8)  # batch, seq_len, dim
K = torch.randn(2, 4, 8)
V = torch.randn(2, 4, 8)

out, weights = scaled_dot_product_attention(Q, K, V)
print(out.shape)
print(weights[0])

小结

注意力机制提供了一种动态的信息路由方式。它不依赖固定卷积窗口,也不需要像 RNN 那样逐步传递状态,因此非常适合并行处理长序列。