自从 2017 年 “Attention Is All You Need”(Vaswani 等)提出 Transformer 架构以来,多头注意力(Multi‑Head Attention, MHA)便成为 NLP 领域最具革命性的创新之一。下面,我们将分步剖析这一机制的原理、实现与实际意义,并通过代码和案例让技术细节更易上手。

一、注意力机制回顾

在深入多头注意力之前,让我们快速回顾一下缩放点积自注意力(Scaled Dot‑Product Self‑Attention)的基本原理:

给定输入序列,将其映射为 Query(Q)Key(K)Value(V) 三组向量:

Attention(Q,K,V)=\text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V

        QKᵀ 计算相似性;
        除以 dk\sqrt{d_k} 控制数值范围;
        softmax 得到注意力权重;
        权重加权 V 得到输出。

这种机制使每个 token 可以“看”到序列内的其他 token,从而捕获上下文语义。Transformer 将其称为 Self-Attention

二、为什么要用“多头”?

多头注意力的核心动机有以下几点:

        多视角观察:不同 head 可以捕捉不同类型的语义关系(例如主谓一致、定语修饰、时间顺序等)。

        子空间优势:每个头在 Q/K/V 的不同子空间操作,可提高表达能力

        并行计算:多个头同时计算,自然适配现代硬件并行处理,大幅提升效率

三、多头注意力的数学推导

设输入维度为 d_model,头数为 h,每个头的维度为 d_k = d_model / h。整个流程如下:

  1. 输入 X ∈ R^{L×d_model}

  2. 线性映射得到 Q, K, V 并 reshape 为 (h, L, d_k)

  3. 每个 head 独立计算注意力;

  4. 将 h 个 Head 的输出拼接为 (L, d_model)

  5. 最后一层线性映射得到最终输出。

简要公式:

head_i = Attention(XW_i^Q, XW_i^K, XW_i^V)
MHA(X) = Concat(head_i) · W^O

图示详解见下:


【插图:Q/K/V 分 head → 并行 Attention → 拼接 → 线性变换】。

四、代码 实现

1.示例

下面带你手写一个简易版 MHA,以加深理解:

import torch
import torch.nn as nn
import torch.nn.functional as F

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

        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):
        B, L, D = x.size()
        # (B, L, D) → (B, h, L, d_k)
        q = self.w_q(x).view(B, L, self.h, self.d_k).transpose(1,2)
        k = self.w_k(x).view(B, L, self.h, self.d_k).transpose(1,2)
        v = self.w_v(x).view(B, L, self.h, self.d_k).transpose(1,2)

        # Scaled dot-product attention
        scores = torch.matmul(q, k.transpose(-2,-1)) / (self.d_k ** 0.5)
        if mask is not None:
            scores = scores.masked_fill(mask == 0, -1e9)
        weights = F.softmax(scores, dim=-1)
        context = torch.matmul(weights, v)  # (B, h, L, d_k)

        # 拼接 + 线性层
        context = context.transpose(1,2).contiguous().view(B, L, D)
        out = self.w_o(context)
        return out, weights

测试运行

x = torch.rand(2, 5, 512)  # batch=2, seq_len=5
mha = MultiHeadAttention(512,8)
o, w = mha(x)
print(o.shape, w.shape)  # (2,5,512), (2,8,5,5)
  • o: 模型输出

  • w: 每个 head 的注意力权重,可用于可视化理解。

2.案例代码

# 多头注意力机制
import copy


def clones(module, N):
    return nn.ModuleList([copy.deepcopy(module) for _ in range(N)])


class MultiHeadedAttention(nn.Module):
    def __init__(self, head, embedding_dim, dropout=0.1):
        super(MultiHeadedAttention, self).__init__()
        assert embedding_dim%head== 0
        self.d_k = embedding_dim // head
        self.head = head
        self.linears = clones(nn.Linear(embedding_dim, embedding_dim), 4)
        self.attn = None
        self.dropout = nn.Dropout(p=dropout)

    def forward(self, query, key, value,  mask=None):

        if mask is not None:
            mask = mask.unsqueeze(0)
        # print('multmaskshape===', mask.shape) #multmaskshape=== torch.Size([1, 8, 4, 4])
        batch_size = query.size(0)
        # view中的四个参数的意义
        # batch_size: 批次的样本数量
        # -1这个位置应该是: 每个句子的长度
        # self.head*self.d_k应该是embedding的维度, 这里把词嵌入的维度分到了每个头中, 即每个头中分到了词的部分维度的特征
        # query, key, value形状torch.Size([2, 8, 4, 64])
        query, key, value = [model(x).view(batch_size, -1,  self.head, self.d_k).transpose(1, 2) for model, x in zip(self.linears, (query, key, value))]
        # query, key, value = [model(x) for model, x in zip(self.linears, (query, key, value))]
        # print('-=-=', query.shape)
        # print('-=-=', key.shape)
        # print('-=-=', value.shape)
        '''
        -=-= torch.Size([2, 4, 512])
        -=-= torch.Size([2, 4, 512])
        -=-= torch.Size([2, 4, 512])
        '''
        # 所以mask的形状 torch.Size([1, 8, 4, 4])  这里的所有参数都是4维度的   进过dropout的也是4维度的
        x, self.attn = attention(query, key, value, mask=mask, dropout=self.dropout)
        # contiguous解释:https://zhuanlan.zhihu.com/p/64551412
        # 这里相当于图中concat过程
        x = x.transpose(1, 2).contiguous().view(batch_size, -1, self.head*self.d_k)

        return self.linears[-1](x)

五、可视化 & 案例解析

 案例:句子 “The cat sat on the mat.”

不同 head 会关注不同类型的关联:

  • head1:冠词与名词关系 (“The”→“cat”,“the”→“mat”);

  • head2:动词修饰 (“sat”→“on”, “on”→“mat”);

  • head3:跨距远距离关注 (“cat”→“mat”);

六、多头注意力在 Transformer 中的位置

Transformer Encoder 中,每一层包括:

x → MHA(x,x,x) → 残差 + LayerNorm → FFN → 残差 + LayerNorm

MHA 的输出为 FFN 输入,而残差结构确保原始位置信息得以保留。

七、多头数 h 的选取与实践

  • 通常 h=8 或 16;

  • 增加 h 提升表达能力,但也会分割维度,导致每 head 可用特征减少;

  • 实际效果 h 与 d_model、数据量、任务相关,是个超参。

八、端到端示例:将 MHA 应用于分类任务

class TransformerClassifier(nn.Module):
    def __init__(self, d_model=512, heads=8, num_layers=2, num_classes=10):
        super().__init__()
        self.embed = nn.Embedding(10000, d_model)
        self.pe = PositionalEncoding(d_model)
        self.layers = nn.ModuleList([
            EncoderBlock(d_model, heads) for _ in range(num_layers)
        ])
        self.fc = nn.Linear(d_model, num_classes)

    def forward(self, x):
        x = self.embed(x)
        x = self.pe(x)
        for layer in self.layers:
            x = layer(x)
        out = x.mean(dim=1)
        return self.fc(out)

说明EncoderBlock 包含 MHA + FFN + LayerNorm + 残差。

九、为什么 MHA 能更懂“语言”

  • 捕获复杂语义模式:多头允许关注时序、语法、实体等多种联系;

  • 并行训练:加速模型训练,摆脱 RNN 串行限制;

  • 强泛化能力:适配翻译、分类、摘要、视觉等多种任务。


十、总结与展望

🔸 核心思想一览

特性说明
Q/K/V 子空间将原始特征投射到多个子空间,学习多样关注模式
并行头多头并行计算,效率高、表达丰富
残差整合最终拼接 + 映射,将多路信息融合输出

🔸 应用拓展

  • NLP:翻译、摘要、问答等;

  • CV:Vision Transformer 的核心模块;

  • 多模态:支持文本-图像、音频-文本等交互任务。

🔸 未来趋势

  • 头剪枝:去除无用头以减小模型;

  • 可解释性增强:深入研究各 head 的关注模式;

  • 高效 Transformer:稀疏注意力、更少资源消耗设计。

结语

多头注意力机制是 Transformer 成为主流 AI 架构的关键,它赋予模型多角度理解数据的能力,同时兼顾并行效率。本文从原理、代码、案例、应用层面深入剖析,帮助你掌握这颗 AI 大脑中的“聚焦引擎”。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐