AI大模型03-揭秘 Transformer 内部“团队协作“:多头注意力与因果掩码的真相

多头注意力与因果掩码:Transformer 内部的"团队协作"

你是否遇到过这种情况:刚搞懂自注意力里的 Q/K/V,紧接着就看到"多头注意力"——为什么要多个头?一个头不够用吗?再看到"因果掩码"更懵:模型生成时到底在防什么?网上搜到的资料要么只说多头是"并行算好几遍",要么把掩码一句话带过。本文将从"团队协作"视角拆解多头注意力与因果掩码:为什么需要 h 个头并行、它们如何分工、掩码矩阵如何保证自回归的因果性,并给出可运行的 PyTorch 代码。

核心概念:先讲"是什么"

1.1 一句话:多头注意力是什么

多头注意力(Multi-Head Attention)就是把自注意力复制 h 份,并行计算,再把结果拼起来。每一份都有自己独立的 Q/K/V 线性投影,因此可以关注不同类型的语义关系。

在原版 Transformer 中,h 通常取 8 或 16;在 GPT 等大模型中,h 可以达到几十甚至上百(比如 GPT-3 的 96 层、每层 96 个头)。多头注意力是整个 Transformer 的"主力部队",配合前馈网络、残差连接和层归一化,堆叠成数十到数百层的深度网络。

回忆上一篇的基础:单头自注意力的公式是 Attention(Q,K,V)=softmax(QKT/dk)VAttention(Q,K,V)=softmax(QK^T/\sqrt{d_k})VAttention(Q,K,V)=softmax(QKT/dk​​)V。多头注意力在此基础上做了两件事——并行做 h 遍、最后拼起来

MultiHead(Q,K,V)=Concat(head1,...,headh)WOMultiHead(Q,K,V) = Concat(head_1, ..., head_h) W^OMultiHead(Q,K,V)=Concat(head1​,...,headh​)WO

headi=Attention(QWiQ,KWiK,VWiV)head_i = Attention(QW_i^Q, KW_i^K, VW_i^V)headi​=Attention(QWiQ​,KWiK​,VWiV​)

其中每个 headihead_iheadi​ 使用自己的一组投影矩阵 WiQW_i^QWiQ​、WiKW_i^KWiK​、WiVW_i^VWiV​。这是论文中最核心的公式之一,也是本篇要拆解的主角。

1.2 为什么"一个头"不够:单一子空间限制

先回答那个最朴素的问题:单头注意力明明已经能"看全场"了,为什么还要多头?

关键在于"看"的方式太单一。单头注意力把 Q/K/V 投影到同一个子空间里,只能捕捉一种相似度度量。想象一个团队开会,只有一个负责人:他会关注谁的发言,完全取决于他自己的兴趣。如果这个负责人只关心技术细节,那产品、运营、市场的发言可能全被他当背景噪音过滤掉——会议开完,信息严重偏科。

同样,语言中的关系是多种多样的:有的词对之间是语法依存("的"修饰前面的名词),有的是长距离指代(“他"指向前文的"小明”),有的是语义相关(“苹果"和"公司”),有的只是位置邻近。一个头只有一组 Q/K/V 投影,只能学会其中一种匹配模式。

graph TD
    subgraph 单头: 一种视角
        A1[一组 Q/K/V 投影] --> B1[只学一种关系模式]
        B1 --> C1[例如只关注语法依存<br/>忽略指代和语义关联]
    end
    subgraph 多头: 多种视角
        A2[h组独立 Q/K/V 投影] --> B2[头1: 局部句法]
        A2 --> B3[头2: 长距离指代]
        A2 --> B4[头3: 语义相关]
        A2 --> B5[头h: 其他关系]
    end

图注:单头 vs 多头——单头只有一种"视角",多头并行多种"视角",各自建模不同的语义关系。

这正是多头注意力存在的意义:让不同的头学会不同的"看问题角度",从而在同一层内同时捕捉多种关系。语言太复杂,一个视角装不下。

1.3 多头注意力:一支团队,各司其职

把多头注意力想象成一支项目团队:h 个头就是 h 名成员,每个人都领到一份相同的任务简报(输入),但各自带着不同的"职业背景"(独立投影矩阵),从不同角度分析问题,最后把分析结果汇总给项目经理(输出投影 W^O)。

flowchart LR
    X[输入 X] --> S[拆分到 h 个头]
    S --> H1[头1 投影 Q1K1V1]
    S --> H2[头2 投影 Q2K2V2]
    S --> H3[头3 投影 Q3K3V3]
    S --> H4[... ...]
    S --> Hh[头h 投影 QhKhVh]
    H1 --> C[Concat 拼接]
    H2 --> C
    H3 --> C
    H4 --> C
    Hh --> C
    C --> O[线性投影 W^O]
    O --> Z[输出]

图注:多头注意力结构——输入并行进入 h 个头,各自投影、各自计算注意力,输出拼接后经 W^O 线性投影,形成最终结果。

有意思的是,研究者在训练好的模型上观察到的"头分工"非常清晰:有的头专门捕捉局部语法关系(比如形容词修饰名词),有的头负责追踪长距离指代(把"他"指向"小明"),有的头甚至表现出跨语言的对应关系。团队分工不是设计出来的,而是训练中自然涌现的——模型发现这样分工效率最高,就自己长成了这样。

原理拆解:把技术深度锚点讲透

2.1 多头拼接与线性投影:h 个头怎么合体

多头注意力的实现细节值得仔细看,因为这里有三个关键的维度变化。

第一步,拆分。设输入维度是 d_model,头数是 h,则每个头的维度是 dk=dmodel/hd_k = d_model / hdk​=dm​odel/h。实践中一般要求 d_model 能被 h 整除。输入向量通过一组"总投影矩阵"映射到整个 d_model 空间,再按 h 份切块,就得到各头的 Q、K、V。

第二步,并行计算。每个头独立执行自注意力,输出形状为 (batch,seq_len,dk)(batch, seq\_len, d_k)(batch,seq_len,dk​)。

第三步,拼接 + 输出投影。把 h 个头的输出沿特征维度拼接,恢复成 (batch,seq_len,dmodel)(batch, seq\_len, d_model)(batch,seq_len,dm​odel),再经过一个输出投影矩阵 WOW^OWO 融合。

维度变化追踪 (d_model=8, h=2, d_k=4):
输入 x:       (batch, 5, 8)        5 个 token,每个 8 维
拆分:         8 维 → 2 个头 × 4 维
头1 输出:     (batch, 5, 4)
头2 输出:     (batch, 5, 4)
拼接:         (batch, 5, 8)        两个头首尾相接
W^O 投影:     (batch, 5, 8)        融合多头信息,维度不变

图注:维度追踪——输入 8 维拆成 2 个头各 4 维,计算后拼接回 8 维,再经 W^O 投影,输出维度与输入一致。

为什么要再经过一次 WOW^OWO 投影,而不是直接把拼接结果当作输出?因为直接拼接只是"把信息堆在一起",而 W^O 允许模型学习"如何组合这些头的信息"——哪些头的信息重要、哪些需要加权、哪些需要交叉融合。少了这一步,多头就真的只是"算几遍然后摞起来"了。

还有一个容易混淆的细节:论文原文用的是每组头独立的 WiQW_i^QWiQ​、WiKW_i^KWiK​、WiVW_i^VWiV​,但主流实现(包括 PyTorch 的 nn.MultiheadAttention)通常用一个大的投影矩阵一次映射整个 d_model,再切块分成 h 个头。两种写法在数学上等价,但后者在 GPU 上更高效——一次大矩阵乘法比 h 次小矩阵乘法更划算。

顺带算一笔账:多头注意力的总参数量和单头完全一样。设 d_model=512、h=8,单头的 Q/K/V 三组投影是 512×512512 \times 512512×512,共 3×512² ≈ 78.6 万参数;多头拆成 8 组 512×64512 \times 64512×64 的投影,总参数量同样是 8×3×512×64≈78.68 \times 3 \times 512 \times 64 ≈ 78.68×3×512×64≈78.6 万——一分不多,一分不少。所以多头并不是"花了 8 倍钱雇了 8 个人",而是"用同样的预算雇了 8 个专才,每人只干自己擅长的一小块"。参数量不变,表达力却因为"多个视角"而提升,这正是多头设计最精妙的地方。

2.2 头部分工:句法头与指代头

多头注意力的"分工"到底长什么样?研究人员做了大量可视化实验,结论非常有趣。

  • 局部句法头:注意力集中在本位置附近的 1~3 个词上。比如在"红色的苹果"中,“红色"这个词的头会强烈关注"苹果”——这是一种修饰关系,属于典型的局部语法依存。
  • 长距离指代头:注意力跨越很长的距离。比如在"小明昨天把钥匙落在办公室,他今天一早就回去找了"中,"他"这个位置的头会跨越十几个词,把注意力集中到"小明"上。
  • 语义相关头:关注词义相关的词对,即使它们位置很远、语法上也没有直接关系。
  • 特殊符号头:某些头专门关注句首标记、标点符号等特殊位置,帮助模型建立全局结构感。
graph LR
    subgraph 句法头: 看局部
        L1[红色的 苹果 很甜] --> L2[红色 → 苹果 权重高]
    end
    subgraph 指代头: 看远方
        R1[小明 昨天把钥匙落在办公室...他 今天一早就回去找了] --> R2[他 → 小明 权重高]
    end

图注:两种典型的头分工——句法头聚焦相邻词的修饰关系,指代头跨越长距离建立指代关联。

这种分工的意义在于:模型不需要在"局部语法"和"长距离指代"之间做取舍。单头注意力被迫用一种权衡(通常是局部优先,因为局部信息更稳定),而多头可以"两头都要"——句法头负责局部,指代头负责远方,互不干扰,各得其所。

2.3 残差连接与层归一化:深度网络的"电梯与稳压器"

多头注意力不是孤军奋战。在真实 Transformer 中,每个子层(注意力或前馈网络)外面都套着两个关键的"保险装置"——残差连接(Residual Connection)和层归一化(Layer Normalization)

残差连接的作用是给梯度开一条"直达电梯":output=x+Sublayer(x)output = x + Sublayer(x)output=x+Sublayer(x)。没有它,几十层的网络梯度很容易在反向传播中消失或爆炸;有了它,梯度可以绕过中间层直接传回输入,深层网络才能稳定训练。数学上看,反向传播时求导会得到 1+∂Sublayer(x)∂x1 + \frac{\partial Sublayer(x)}{\partial x}1+∂x∂Sublayer(x)​——那个常数 1 保证了梯度在最坏情况下也有一条"保底通道"不会被连续乘法磨没。这就像老式居民楼装了电梯:楼层从 6 层盖到 60 层,如果还只靠楼梯,谁爬得动?残差连接就是给信息流和梯度流装的那部电梯。

层归一化则是"稳压器":对每个 Token 的特征向量做归一化(减均值、除标准差),让每层的输入分布保持稳定,防止数值随着层数增加无限漂移。它和 Batch Normalization 的区别值得留意:BN 在批量维度上做统计,受 batch size 影响且依赖批量统计量,在变长序列和在线推理场景下不方便;LN 则在每个样本的特征维度上独立归一化,不依赖 batch、不依赖序列长度,天然适配 Transformer 这种"一条样本一个变长序列"的结构。这也是 LN 成为 Transformer 标配的工程原因。

flowchart TD
    X[输入 x] --> A[多头注意力]
    X --> R1[残差连接: x + Attn x]
    A --> R1
    R1 --> N1[层归一化]
    N1 --> F[前馈网络 FFN]
    N1 --> R2[残差连接: + FFN 结果]
    F --> R2
    R2 --> N2[层归一化]
    N2 --> Y[输出]
    style R1 fill:#e8f5e9
    style R2 fill:#e8f5e9
    style N1 fill:#fff3e0
    style N2 fill:#fff3e0

图注:一个 Transformer Block 的标准结构——残差连接(绿色)与层归一化(橙色)包裹着注意力和前馈网络两个子层。

有个常见的混淆点:层归一化做在哪里有两种流派。原版论文用 Post-LN(归一化在残差之后),而 GPT 等现代大模型普遍用 Pre-LN(归一化在残差之前,如 x = x + Sublayer(Norm(x)))。Pre-LN 训练更稳定,是工程实践的主流。本文代码采用 Pre-LN 风格。

2.4 因果掩码:生成时不能"偷看未来"

现在进入本篇的另一半重头戏:因果掩码(Causal Mask / Masked Attention)

前面讲的自注意力是"双向"的:每个位置都能看到全场所有位置。这在理解任务(如 BERT 的完形填空)中是合理的,但在生成任务中却是致命的——如果模型在生成第 5 个词时能"看到"第 6 个词,那它就不是在生成,而是在抄答案。

**自回归生成(Autoregressive Generation)**的要求是:生成第 t 个词时,只能依据前 t-1 个词(以及当前的起点),绝不能使用任何未来的信息。这就是"因果"(Causal)的含义——原因在前,结果在后,未来的信息不能影响过去。

实现方式非常巧妙:在 Softmax 之前,把注意力分数矩阵的上三角区域全部填充为 -inf,下三角和主对角线保持原值:

因果掩码矩阵 (5×5, 1=可见, 0=屏蔽):
        位置1  位置2  位置3  位置4  位置5
位置1  [  1     0     0     0     0 ]
位置2  [  1     1     0     0     0 ]
位置3  [  1     1     1     0     0 ]
位置4  [  1     1     1     1     0 ]
位置5  [  1     1     1     1     1 ]
第 t 行: 只能看到前 t 个位置(含自己),未来全部屏蔽

图注:因果掩码矩阵——下三角(含主对角线)为 1 表示可见,上三角为 0 表示屏蔽,保证每个位置只能关注自己及之前的位置。

为什么是 -inf 而不是 0?这是本篇最重要的一个"工程细节"。Softmax 会对输入做指数运算:softmax(x)i=exi/∑jexjsoftmax(x)_i = e^{x_i}/\sum_j e^{x_j}softmax(x)i​=exi​/∑j​exj​。如果把屏蔽位置填 0,它仍然会获得 e0=1e^0=1e0=1 的权重,根本没有被屏蔽;而填 -inf 后,e−inf=0e^{-inf}=0e−inf=0,该位置的权重严格为 0——被屏蔽的位置彻底"消失",且不参与归一化分母,其余位置的注意力权重会自动重新分配、加起来仍等于 1。

示例: 位置3 的分数 [0.5, 0.3, 0.9, 1.2, -inf]
Softmax 后:     [0.14, 0.11, 0.20, 0.55, 0.00]
位置5 被 -inf 屏蔽 → 权重严格为 0, 前四个位置重归一化

图注:-inf 屏蔽效果——被屏蔽位置的 Softmax 权重严格为 0,且不影响其他位置的相对比例。

如果把这个掩码作用到整个批量上,矩阵乘法和 Softmax 可以一次性完成,效率不受影响——掩码只是给分数矩阵"打洞",不增加计算量。

这里藏着一个新手最容易忽视的细节:掩码在训练和推理时的用法不一样。推理时,模型确实只生成一个词、只看过去的词;但训练时,如果也逐词生成再对比答案,那一个句子训练 n 个词就要前向传播 n 次,慢得没法用。所以训练用的是"教师强制"(Teacher Forcing)技巧:把整个目标句子一次性输入,每个位置的预测同时并行算出,用因果掩码保证"第 t 个位置的预测只看前 t 个位置"。这样一次前向传播就能算出整句所有位置的下一个词预测,损失函数对所有位置的预测统一计算梯度。掩码在这里扮演的角色,就是在"并行训练"和"自回归推理"之间搭桥:训练时并行,行为上却严格等价于逐词生成。

再补充一个容易误解的点:因果掩码屏蔽的是"未来位置的信息",但每个位置仍然能看到自己左侧的全部历史——所以注意力权重不是只集中在最后一个词上,而是可以分布在从开头到当前的所有 Token 上,各按相关性分配。这意味着长距离指代在生成模型里同样成立:“他"在生成时,可以回头把大部分注意力给到几十个词之前的"小明”——只要掩码允许"看到过去",多远都没关系。

2.5 自回归生成:GPT 的"逐词预测"机制

把因果掩码装进解码器,就得到了 GPT 一族的基本结构。自回归生成的流程是"一次一个词"的循环

  1. 输入当前已有的词序列,经过多层带因果掩码的 Transformer Block;
  2. 取最后一个位置的输出,通过词汇表映射层得到"下一个词的概率分布";
  3. 按概率采样(或取最大概率)选出下一个词;
  4. 把新词拼到序列末尾,重复第 1 步,直到生成结束标记或达到长度上限。
flowchart LR
    A[已有序列: 今天天气] --> B[带因果掩码的<br/>Transformer Block]
    B --> C[最后一个位置的输出]
    C --> D[词汇表概率分布]
    D --> E[采样下一个词: 很]
    E --> F[新序列: 今天天气很]
    F -.重复循环.-> A

图注:自回归生成循环——每轮只预测一个词,新词拼回输入,如此往复,整个序列像"吐字"一样逐词生成。

这就是为什么 ChatGPT 的回复总是"逐字蹦出来"而不是一次性整段输出:架构上它只能自回归地一个 Token 一个 Token 生成。也正因为每一步只能看过去,生成的连贯性完全依赖模型在训练中学到的"下一个词预测"能力——它从海量语料中学到"在这种上下文里,下一个词大概率是什么",这就是大语言模型最底层的运作逻辑。

最后补一句采样细节,帮你把"概率分布"和"实际生成"连起来。模型每一步输出的是整个词汇表上的概率分布,但我们并不会每次都用概率最大的那个词——那样生成的文本会单调重复(专业叫法是"退化为模板")。实际生成会配合温度参数(temperature)和采样策略:温度调高,分布变平坦,生成更大胆、更发散;温度调低,分布变尖锐,生成更保守、更稳定。而无论怎么采样,每一步的输入都来自上一步真实生成的 Token,而不是概率分布本身——如果上一步生成了错词,后续内容也只能在"错词的地基"上继续生长。这种"走一步看一步、错了也要走下去"的特性,就是自回归生成既强大又偶尔翻车的根源。

实战演练:PyTorch 实现 Multi-Head Attention + 因果掩码

理论到位,动手验证。下面用 PyTorch 实现多头注意力 + 因果掩码,代码可以直接运行。

3.1 完整代码

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


def build_causal_mask(seq_len: int) -> torch.Tensor:
    """生成因果掩码: 下三角为 1(可见), 上三角为 0(屏蔽)"""
    mask = torch.tril(torch.ones(seq_len, seq_len))
    # 形状: (1, 1, seq_len, seq_len),方便广播到 (batch, h, seq_len, seq_len)
    return mask.unsqueeze(0).unsqueeze(0)


class MultiHeadAttention(nn.Module):
    """多头自注意力 + 可选因果掩码"""

    def __init__(self, d_model: int, h: int):
        super().__init__()
        assert d_model % h == 0, "d_model 必须能被 h 整除"
        self.d_model = d_model
        self.h = h
        self.d_k = d_model // h

        # 统一投影: 一次映射整个 d_model,再切块给 h 个头(主流高效实现)
        self.W_q = nn.Linear(d_model, d_model, bias=False)
        self.W_k = nn.Linear(d_model, d_model, bias=False)
        self.W_v = nn.Linear(d_model, d_model, bias=False)
        self.W_o = nn.Linear(d_model, d_model, bias=False)

    def forward(self, x, mask=None):
        """
        x: (batch, seq_len, d_model)
        mask: (1, 1, seq_len, seq_len),None 表示全可见(双向)
        """
        batch, seq_len, _ = x.shape

        # 1. 投影 → (batch, seq_len, d_model) → 拆成 h 个头
        Q = self.W_q(x).view(batch, seq_len, self.h, self.d_k).transpose(1, 2)
        K = self.W_k(x).view(batch, seq_len, self.h, self.d_k).transpose(1, 2)
        V = self.W_v(x).view(batch, seq_len, self.h, self.d_k).transpose(1, 2)
        # 此时 Q/K/V: (batch, h, seq_len, d_k)

        # 2. 点积 + √dk 缩放
        scores = torch.matmul(Q, K.transpose(-2, -1)) / (self.d_k ** 0.5)
        # scores: (batch, h, seq_len, seq_len)

        # 3. 因果掩码: 被屏蔽位置填 -inf,Softmax 后权重为 0
        if mask is not None:
            scores = scores.masked_fill(mask == 0, float("-inf"))

        # 4. Softmax + 加权求和
        attn = F.softmax(scores, dim=-1)
        out = torch.matmul(attn, V)          # (batch, h, seq_len, d_k)

        # 5. 拼接 h 个头: (batch, seq_len, d_model)
        out = out.transpose(1, 2).contiguous().view(batch, seq_len, self.d_model)

        # 6. 输出线性投影 W^O
        return self.W_o(out)


# 演示: 双向 vs 因果掩码的注意力权重差异
if __name__ == "__main__":
    torch.manual_seed(42)
    batch, seq_len, d_model, h = 1, 5, 8, 2
    x = torch.randn(batch, seq_len, d_model)
    mha = MultiHeadAttention(d_model, h)

    # 双向模式(不传 mask)
    out_bi = mha(x)
    print("双向输出形状:", tuple(out_bi.shape))

    # 因果模式(传 mask)
    causal_mask = build_causal_mask(seq_len)
    out_causal = mha(x, mask=causal_mask)
    print("因果输出形状:", tuple(out_causal.shape))
    print("\n因果掩码矩阵:\n", causal_mask.squeeze().numpy())

3.2 运行效果

双向输出形状: (1, 5, 8)
因果输出形状: (1, 5, 8)
因果掩码矩阵:
 [[1. 0. 0. 0. 0.]
  [1. 1. 0. 0. 0.]
  [1. 1. 1. 0. 0.]
  [1. 1. 1. 1. 0.]
  [1. 1. 1. 1. 1.]]

注意看掩码矩阵:下三角全是 1,上三角全是 0,主对角线也是 1(每个位置都能看到自己)。这正是自回归生成要求的"只能看过去"。

你可以做个对照实验验证因果性:取"双向模式"下第 1 行第 3 列的注意力权重(位置 1 对位置 3 的关注),它通常是一个非零值;再取"因果模式"下同一个位置,它严格等于 0——因为 -inf 被 Softmax 归零了。这个对照,就是因果掩码作用的直接证据。

3.3 代码结构图

flowchart TD
    A[MultiHeadAttention] --> B[投影 Wq/Wk/Wv]
    B --> C[view + transpose 拆 h 个头]
    C --> D[点积分数 / √dk]
    D --> E{有 mask?}
    E -->|是| F[masked_fill 填 -inf]
    E -->|否| G[Softmax]
    F --> G
    G --> H[加权求和 @ V]
    H --> I[拼接 + W^O 投影]
    I --> J[输出]

图注:代码结构——投影、拆头、缩放、可选掩码、Softmax、加权、拼接投影,七步与理论一一对应。

3.4 踩坑点总结

踩坑点现象解决方案
d_model 不能被 h 整除运行时形状错误断言检查,或调整 d_model/h 满足整除
transpose 后忘 contiguous()view 报错"非连续内存"拼接前先 .contiguous()
掩码填 0 而非 -inf被屏蔽位置仍有权重float("-inf"),Softmax 后权重才严格为 0
掩码形状不广播报维度不匹配扩到 (1,1,seq_len,seq_len),依赖广播
只看最后一个输出以为"生成"是整体输出自回归要循环逐词预测(见 2.5)

避坑指南与效率技巧

⚠️ 避坑警告一:别把"多头"理解为"模型变强 h 倍"。多头不增加参数量级的收益是"视角多样化",不是"能力翻倍"。论文里单头和多头的效果差距,主要体现在长距离和复杂关系建模上,简单任务上多头优势并不明显。跟人介绍时说"并行多种视角"比说"算得更快更多"准确得多。

⚠️ 避坑警告二:因果掩码一定要在 Softmax 之前、且用 -inf。这是新手最容易踩的坑。掩码顺序错了(先 Softmax 再掩码),被屏蔽位置的权重不会归零;掩码值用了 0 而不是 -inf,屏蔽形同虚设。生成模型训练时"偷看未来",loss 会异常地低,但推理时效果崩盘——典型的"训练一时爽,上线火葬场"。

⚠️ 避坑警告三:不要在所有场景都套因果掩码。只有自回归生成(GPT 类解码器)需要因果掩码;BERT 类编码器的完形填空、句子分类等理解任务,用的是双向注意力(不掩码)。套错了会白白牺牲理解能力。

⚠️ 避坑警告四:用 KV 缓存加速推理时,掩码必须是"动态增长"的。每生成一个新 Token,能看到的历史就多一个位置,掩码矩阵要从 N×N 扩到 (N+1)×(N+1)。如果你把掩码写死成固定大小,第一次生成没问题,第二次就开始形状报错;就算硬凑过去,也大概率把新 Token 的位置屏蔽错了,生成结果会莫名其妙地变差。生产代码里务必每步根据当前序列长度重新构建或切片掩码。

💡 效率技巧一:用 PyTorch 内置 nn.MultiheadAttention 验证自己的实现。把 is_causal=True 传进去,和手写实现对比输出,一行代码就能交叉验证正确性。生产环境也直接用官方封装,性能更优。

💡 效率技巧二:生成时用 KV 缓存,别每步全量重算。自回归每步都重新对前面所有 Token 算注意力是浪费——把已经算好的 K、V 缓存起来,每步只算新 Token,推理速度能提升数倍。PyTorch 和 transformers 都原生支持。

💡 效率技巧三:掩码先构建一次,训练循环里复用。因果掩码只跟序列长度有关,跟 batch 无关——同一个长度下的掩码张量可以在每个 batch 复用,不需要每步重建。把 build_causal_mask 的结果缓存起来,省掉重复的 CPU 到 GPU 拷贝,训练脚本能小有提速;配合半精度训练,效果更佳。

总结与思考

回到开篇的两个困惑。多头注意力解决了"一个视角装不下复杂语言"的问题:h 组独立投影并行计算,让不同的头各司其职——句法头看局部、指代头看远方、语义头看关联,最后通过拼接和 W^O 投影融合成一个丰富的表示。因果掩码解决了"生成时不能偷看未来"的问题:把注意力分数矩阵的上三角填成 -inf,Softmax 后权重严格归零,让每个位置只能依据过去生成未来,这正是自回归模型的因果性根基。

这两者共同构成了大语言模型(GPT 一族)的"发动机":多头保证每层都能从多角度理解上下文,掩码保证生成过程严格单向、逐词推进。配合上一篇讲的自注意力,你已经掌握了 Transformer 最核心的三块拼图;再加上残差连接与层归一化这对"电梯与稳压器",一个完整的解码器 Block 已经在你脑中成型了。

金句:多头注意力让模型"眼观六路",因果掩码让模型"只看过去"——前者决定了它有多会理解,后者决定了它有多会生成。

如果你能回答下面三个问题,说明这篇吃透了:① 为什么掩码用 -inf 而不是 0?② 多头拼接后为什么还要一次 W^O 投影?③ 自回归生成中,第 t 个位置的注意力能看到哪些位置?

互动与转化

互动钩子:你在用生成模型时,有没有遇到过"模型似乎在偷看答案"的诡异现象?或者调试多头注意力时踩过形状不匹配的坑?评论区聊聊你的经历——顺手点个收藏,下次面试被问"为什么用 -inf 做掩码"时,直接翻这篇。

【系列文章预告】 下一篇我将带你拆解【Token 化与位置编码:模型看到的不是"字"】,看文本是如何被切成 Token、模型如何感知词序的,敬请关注。

文末标签

#多头注意力 #Transformer #掩码 #自回归 #大模型 #深度学习 #NLP

评论
成就一亿技术人!
拼手气红包6.0元
还能输入1000个字符
 
 条评论被折叠 查看
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包

打赏作者

每日干货分享

你的鼓励将是我创作的最大动力

¥1 ¥2 ¥4 ¥6 ¥10 ¥20
扫码支付:¥1
获取中
扫码支付

您的余额不足,请更换扫码支付或充值

打赏作者

实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值