多头注意力与因果掩码: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=dmodel/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,dmodel),再经过一个输出投影矩阵 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/∑jexj。如果把屏蔽位置填 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 一族的基本结构。自回归生成的流程是"一次一个词"的循环:
- 输入当前已有的词序列,经过多层带因果掩码的 Transformer Block;
- 取最后一个位置的输出,通过词汇表映射层得到"下一个词的概率分布";
- 按概率采样(或取最大概率)选出下一个词;
- 把新词拼到序列末尾,重复第 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

434

被折叠的 条评论
为什么被折叠?



