注意力机制详解

Transformer 的灵魂,让模型学会”关注重点”。


自注意力(Self-Attention)

核心思想

让序列中的每个元素都能”看到”其他所有元素,并决定应该关注谁。

Q、K、V 三元组

每个词被映射为三个向量:

Q (Query)  - 查询:"我想找什么信息?"
K (Key)    - 键:"我有什么特征可以被匹配?"
V (Value)  - 值:"如果匹配上了,提供什么内容?"

计算流程

步骤 1:计算相似度
        Score = Q × K^T

步骤 2:缩放
        Scaled Score = Score / √d_k

步骤 3:归一化
        Attention Weights = softmax(Scaled Score)

步骤 4:加权求和
        Output = Attention Weights × V

多头注意力(Multi-Head Attention)

为什么需要多头?

单头注意力只能学习一种关注模式,多头可以同时学习多种:

句子:"小明在银行边的长椅上看书"

┌─────────────────────────────────────────────────┐
│  头1(语法专家):主谓宾关系                        │
│    小明 ←──────→ 看书                            │
│                                                  │
│  头2(位置专家):空间关系                          │
│    长椅 ←── 在 ──→ 银行边                         │
│                                                  │
│  头3(修饰专家):定语关系                          │
│    长椅 ←── 的 ──→ 银行边                         │
│                                                  │
│  头4(语义专家):词义消歧                          │
│    银行 + 边 + 长椅 → 大概率是"河岸"               │
└─────────────────────────────────────────────────┘
              ↓ 综合所有头的分析结果 ↓
         得到更全面、更准确的理解

数学表示

# 多头注意力计算
MultiHead(Q, K, V) = Concat(head_1, ..., head_h) × W_O
 
# 每个头
head_i = Attention(Q × W_Q_i, K × W_K_i, V × W_V_i)

代码示例

class MultiHeadAttention:
    def __init__(self, d_model=512, num_heads=8):
        self.num_heads = num_heads
        self.d_k = d_model // num_heads  # 每头维度
 
        # 每个头有独立的投影矩阵
        self.W_Q = Linear(d_model, d_model)
        self.W_K = Linear(d_model, d_model)
        self.W_V = Linear(d_model, d_model)
        self.W_O = Linear(d_model, d_model)
 
    def forward(self, Q, K, V):
        batch_size = Q.size(0)
 
        # 1. 线性投影并拆分成多头
        Q = self.W_Q(Q).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2)
        K = self.W_K(K).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2)
        V = self.W_V(V).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2)
 
        # 2. 计算注意力
        scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
        attn_weights = F.softmax(scores, dim=-1)
        context = torch.matmul(attn_weights, V)
 
        # 3. 合并多头
        context = context.transpose(1, 2).contiguous().view(batch_size, -1, self.num_heads * self.d_k)
 
        # 4. 最终投影
        output = self.W_O(context)
        return output

为什么”多头”比”一个大头”好?

方案问题
一个 512 维的大头所有关系混在一起学,容易互相干扰
8 个 64 维的小头分工明确,各学各的,最后综合

类比

  • 单头 = 一个全科医生看所有病
  • 多头 = 内科、外科、眼科等专家会诊

实际模型配置

模型头数每头维度总维度
BERT-base1264768
GPT-21264768
GPT-39612812288
GPT-4 (推测)120+128+更大

注意力的变体

1. 因果注意力(Causal Attention)

用于生成模型,每个位置只能看到之前的内容:

      看  书  真  好
看    ✓   ✗   ✗   ✗
书    ✓   ✓   ✗   ✗
真    ✓   ✓   ✓   ✗
好    ✓   ✓   ✓   ✓

2. 交叉注意力(Cross Attention)

Q 来自一个序列,K 和 V 来自另一个序列,用于编码器-解码器架构。

3. 稀疏注意力(Sparse Attention)

不是所有位置都互相注意,降低计算复杂度。


相关链接


创建时间: 2024-12-10