多头注意力机制(Multi-head Attention)的演进与应用:从发明到Transformer革命

1次阅读
没有评论

共计 2075 个字符,预计需要花费 6 分钟才能阅读完成。

image.webp

技术溯源

多头注意力机制 (Multi-head Attention) 首次出现在 2017 年 6 月 Google 发表的论文《Attention Is All You Need》中,这篇论文提出了革命性的 Transformer 架构,彻底改变了自然语言处理领域的格局。

多头注意力机制 (Multi-head Attention) 的演进与应用:从发明到 Transformer 革命

传统序列建模主要依赖 RNN 和 CNN,但存在明显缺陷:

  • RNN 难以并行计算,且面临长距离依赖问题
  • CNN 的感受野有限,需要堆叠多层才能捕获全局信息

自注意力机制 (Self-Attention) 通过计算序列中所有位置的关系权重,能够直接建模任意距离的依赖关系。多头注意力则进一步扩展了这一思想。

机制拆解

数学表达

多头注意力的核心公式如下:

$$
\text{MultiHead}(Q,K,V) = \text{Concat}(head_1,…,head_h)W^O
$$

其中每个注意力头 (head) 的计算为:

$$
head_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)
$$

架构流程

  1. 输入向量通过线性变换分成 Query、Key、Value
  2. 拆分为 h 个头,每个头独立计算注意力
  3. 所有头的输出拼接后通过线性变换得到最终结果
  4. 维度关系:每个头的维度 d_k = d_model / h

关键超参数

  • 头数 h 通常取 8 的倍数
  • d_model 必须能被 h 整除以保证维度一致性
  • 实践中常用配置:d_model=512, h=8

PyTorch 实战

import torch
import torch.nn as nn
import math

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model=512, h=8):
        super().__init__()
        assert d_model % h == 0, "d_model must be divisible by h"

        self.d_k = d_model // h
        self.h = h

        # 线性变换层
        self.W_q = nn.Linear(d_model, d_model)
        self.W_k = nn.Linear(d_model, d_model)
        self.W_v = nn.Linear(d_model, d_model)
        self.W_o = nn.Linear(d_model, d_model)

    def forward(self, x, mask=None):
        # x: [batch_size, seq_len, d_model]
        batch_size = x.size(0)

        # 线性变换并分头 [batch_size, seq_len, h, d_k]
        Q = self.W_q(x).view(batch_size, -1, self.h, self.d_k)
        K = self.W_k(x).view(batch_size, -1, self.h, self.d_k)
        V = self.W_v(x).view(batch_size, -1, self.h, self.d_k)

        # 转置为 [batch_size, h, seq_len, d_k]
        Q = Q.transpose(1, 2)
        K = K.transpose(1, 2)
        V = V.transpose(1, 2)

        # 计算注意力分数 [batch_size, h, seq_len, seq_len]
        scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)

        # 应用 mask(解码器用)if mask is not None:
            scores = scores.masked_fill(mask == 0, -1e9)

        # 计算注意力权重
        attention = torch.softmax(scores, dim=-1)

        # 加权求和 [batch_size, h, seq_len, d_k]
        context = torch.matmul(attention, V)

        # 拼接多头输出 [batch_size, seq_len, d_model]
        context = context.transpose(1, 2).contiguous().view(batch_size, -1, self.h * self.d_k)

        return self.W_o(context)

生产建议

头数选择

  • 经验公式:h ≈ d_model / 64
  • 常见配置:
  • BERT: d_model=768, h=12
  • GPT-3: d_model=12288, h=96

显存优化

  • 多头注意力显存占用与 batch_size * h * seq_len^2 成正比
  • 长序列场景可考虑:
  • 减少头数
  • 使用内存高效的注意力实现

可视化工具

  • BertViz:直观展示注意力模式
  • TensorBoard:监控注意力权重分布

延伸思考

Vision Transformer 应用

  • 将图像切分为 patch 作为输入序列
  • 位置编码适应 2D 空间关系
  • 注意力头可捕获不同视觉模式

开放问题

  • 头数增加可能带来冗余
  • Attention Head Dropout 可提升模型鲁棒性
  • 动态头数分配是潜在优化方向

多头注意力机制已经成为现代深度学习架构的基石组件,理解其原理和实现细节对于设计和优化 Transformer 模型至关重要。随着研究的深入,这一技术仍在不断演进,展现出更广阔的应用前景。

正文完
 0
评论(没有评论)