现代Transformer精读-03:自注意力机制精读
本文精读
model_minimind.py的Attention类(91–134 行)与repeat_kv(86–89 行)。
0. 整块代码:Attention 类全貌
正文每一节都对应其中的几行
# model_minimind.py 86–134 行
def repeat_kv(x: torch.Tensor, n_rep: int) -> torch.Tensor:
bs, slen, num_key_value_heads, head_dim = x.shape
if n_rep == 1: return x
return (x[:, :, :, None, :].expand(bs, slen, num_key_value_heads, n_rep, head_dim)
.reshape(bs, slen, num_key_value_heads * n_rep, head_dim))
class Attention(nn.Module):
def __init__(self, config: MiniMindConfig):
super().__init__()
self.num_key_value_heads = config.num_attention_heads if config.num_key_value_heads is None else config.num_key_value_heads
self.n_local_heads = config.num_attention_heads # 8 个 Q 头
self.n_local_kv_heads = self.num_key_value_heads # 4 组 KV 头
self.n_rep = self.n_local_heads // self.n_local_kv_heads # 2:每个 KV 头服务 2 个 Q 头
self.head_dim = config.head_dim # 96
self.is_causal = True
self.q_proj = nn.Linear(config.hidden_size, config.num_attention_heads * self.head_dim, bias=False)
self.k_proj = nn.Linear(config.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)
self.v_proj = nn.Linear(config.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)
self.o_proj = nn.Linear(config.num_attention_heads * self.head_dim, config.hidden_size, bias=False)
self.q_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps)
self.k_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps)
self.attn_dropout = nn.Dropout(config.dropout)
self.resid_dropout = nn.Dropout(config.dropout)
self.dropout = config.dropout
self.flash = hasattr(torch.nn.functional, 'scaled_dot_product_attention') and config.flash_attn
def forward(self, x, position_embeddings, past_key_value=None, use_cache=False, attention_mask=None):
bsz, seq_len, _ = x.shape
xq, xk, xv = self.q_proj(x), self.k_proj(x), self.v_proj(x) # ① 投影
xq = xq.view(bsz, seq_len, self.n_local_heads, self.head_dim) # ② 切头
xk = xk.view(bsz, seq_len, self.n_local_kv_heads, self.head_dim)
xv = xv.view(bsz, seq_len, self.n_local_kv_heads, self.head_dim)
xq, xk = self.q_norm(xq), self.k_norm(xk) # ③ QK-Norm
cos, sin = position_embeddings
xq, xk = apply_rotary_pos_emb(xq, xk, cos, sin) # ④ RoPE(02 篇)
if past_key_value is not None: # ⑤ KV 拼接(推理续写)
xk = torch.cat([past_key_value[0], xk], dim=1)
xv = torch.cat([past_key_value[1], xv], dim=1)
past_kv = (xk, xv) if use_cache else None
xq, xk, xv = (xq.transpose(1, 2), # ⑥ 形状:B,heads,S,d
repeat_kv(xk, self.n_rep).transpose(1, 2),
repeat_kv(xv, self.n_rep).transpose(1, 2))
if self.flash and (seq_len > 1) and (not self.is_causal or past_key_value is None) \
and (attention_mask is None or torch.all(attention_mask == 1)): # ⑦ Flash 路径
output = F.scaled_dot_product_attention(xq, xk, xv,
dropout_p=self.dropout if self.training else 0.0,
is_causal=self.is_causal)
else: # ⑧ 手动回退路径
scores = (xq @ xk.transpose(-2, -1)) / math.sqrt(self.head_dim)
if self.is_causal:
scores[:, :, :, -seq_len:] += torch.full((seq_len, seq_len),
float( -inf ), device=scores.device).triu(1) # 上三角 -inf
if attention_mask is not None:
scores += (1.0 - attention_mask.unsqueeze(1).unsqueeze(2)) * -1e9
output = self.attn_dropout(F.softmax(scores.float(), dim=-1).type_as(xq)) @ xv
output = output.transpose(1, 2).reshape(bsz, seq_len, -1) # ⑨ 收回 768
output = self.resid_dropout(self.o_proj(output))
return output, past_kv
1. 缩放点积注意力的数学
1.1 Q / K / V 的语义
自注意力的输入只有一个:残差流里的 x([B, S, 768],§0 的 ①)。它被三个投影(q_proj / k_proj / v_proj)各映射一份,语义上是三种角色:
- Q(Query 查询):当前 token 想找出谁跟我相关 每人提一个问题;
- K(Key 键):每个 token 的 被查询标识 每个位置回答我是谁,供所有问题比对;
- V(Value 值):每个 token 真正携带的内容,被查询结果挑选后加权输出。
写成矩阵形式(§0 的 ⑧):
从数学上理解,其本质意义是通过点积计算各个token之间的语义相关性,然后将获得的权重与各个token携带的内容相乘进行混合
1.2 为什么除以
的每一项是 维内积。若 q、k 的各分量均值 0、方差 1,则:
内积的方差随维度 线性增长。除以 后方差回到 1让打分尺度与维度无关。如果不除:
- 越大,score 的绝对值越大,softmax 输入进入 的饱和区注意力要么近乎 one-hot(只盯一个 token),要么梯度极其平缓(几乎学不动);
- 除以 后输入保持在 ,softmax 的梯度和温度稳定。
1.3 行归一化
注意 softmax 是按行做的:对固定的 ,。每一行就是 当前 token 把自己的注意力在全部历史位置之间分配 。
- 与 RoPE 的关系:分值里其实带着 02 篇的相位 含 ,所以注意力分配天然被相对距离调节(02-RoPE位置编码与YaRN外推);
- 与掩码的联动:未来位置、padding 位置的 score 被压成 ,softmax 后权重为 0.
1.4 复杂度: 注意力的复杂度是O(S²)
打分矩阵的形状是 [B, H, S, S]每个位置对都要算一次内积。8 层、每头 :
序列长度一翻倍,注意力算力涨 4 倍。这也正是 Flash Attention(第 5 节)和 KV Cache(第 6 节)存在的理由。
2. 多头 → GQA → repeat_kv
2.1 为什么要 多头
一个注意力头只有一种相似度度量,它只能学一种谁跟谁相关 的模式。多头(Multi-Head Attention)的本质:并行跑多个独立注意力头,每个头负责一种关系模式(有的关注局部语法、有的关注长程指代、有的关注同一个词的不同含义)。
把 768 维切成 8 份 96 维(§0 的 ②,只是 view,不增加参数),最后 o_proj 把 8 个头的输出拼回 768(§0 的 ⑨)。
2.2 历史:MHA → MQA → GQA
多头 之下还有个问题:KV 要不要也多头?历史上有三种做法:
| 方案 | Q 头 | KV 头 | 代表 |
|---|---|---|---|
| MHA(Multi-Head) | 8 | 8 | 2017 原版 |
| MQA(Multi-Query) | 8 | 1 | PaLM(省到极致,质量有损) |
| GQA(Grouped-Query) | 8 | 4(分 2 组) | Llama 2/3、Qwen、MiniMind |
GQA 的动机:K、V 头少 → ① K/V 投影参数减半② KV Cache 减半推理时逐 token 缓存的正是 K、V,这是推理显存的大头。采用分组注意力,质量几乎无损,显存省去一半(Llama2-70B 激进到 64:8)。
2.3 repeat_kv:4 组 KV 如何服务 8 个 Q
def repeat_kv(x, n_rep): # n_rep = 8 // 4 = 2
bs, slen, num_key_value_heads, head_dim = x.shape # [B, S, 4, 96]
if n_rep == 1: return x
return (x[:, :, :, None, :] # ① 插入新维 [B,S,4,1,96]
.expand(bs, slen, num_key_value_heads, n_rep, head_dim) # ② 广播 [B,S,4,2,96]
.reshape(bs, slen, num_key_value_heads * n_rep, head_dim)) # ③ 摊平 [B,S,8,96]
三步的形状与语义:
| 步骤 | 形状 | 形状含义 |
|---|---|---|
| 输入 | [B, S, 4, 96] | 4 组 KV |
① x[…, None, :] | [B,S,4,1,96] | 在头和特征之间插入一维 |
② expand(...) | [B,S,4,2,96] 视口 | 广播复制:不分配新内存(只改 strides),每组 KV 变成 2 份 |
③ reshape(...) | [B,S,8,96] | reshape分成 8 组 |
于是 Q 头与 KV 组的配对关系是:
同一组 KV 被两个 Q 头共享,但打分各自独立
transpose(1,2) 把 头维 提到第 2 维,得到 [B, H, S, d]一次矩阵乘同时算所有头(xq @ xkᵀ 在 [B,H,S,d]×[B,H,d,S] 上批量完成)。
2.5 呼应 01 篇的账
GQA 在 MiniMind 的具体收益(8:4):
- 参数:k/v 投影 而非 ,每层省 ,8 层省 ≈ 4.7M(约占 7%,见 01-模型骨架与前向回路);
- KV Cache:缓存的 KV 减半
3. QK-Norm:RoPE 之前的稳定化
3.1 代码与位置
# __init__(104–105 行),复用 01 篇的 RMSNorm,但归一化对象是 head_dim
self.q_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps) # 96 维
self.k_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps)
# forward(117–119 行):在 RoPE 注入 **之前**
xq, xk = self.q_norm(xq), self.k_norm(xk)
cos, sin = position_embeddings
xq, xk = apply_rotary_pos_emb(xq, xk, cos, sin)
它把 §0 流水线里 ③ QK-Norm ④ RoPE 的顺序钉死了:先稳定内容,再注入位置。归一化的对象不是整条残差流(768 维),而是每个头自己的 96 维向量RMSNorm 沿最后一维做,所以 8 个 Q 头、4 个 K/V 头各自独立归一。
3.2 注意力 logits 的尺度漂移
[[#1.2 为什么除以 ]] 说注意力对 logits 的尺度敏感(softmax 饱和)。然而在Pre-Norm 时代无人在意 q、k 的尺度
- 回顾 01-模型骨架与前向回路:Pre-Norm 的代价是残差流尺度随深度缓慢增长(final RMSNorm 只在最外层兜底);
- q、k 是
q_proj(x)/k_proj(x)的投影x的尺度在长, 训练中也在变,、 的范数会无界漂移; - 于是 的绝对值跟着漂,softmax 时而被烫得接近 one-hot、时而冷得梯度趋平,严重时某些头整个失活(logits 全被压扁)。
QK-Norm 的做法很直接:在打分之前,把每个头的 q、k 各自的尺度归一打分退回到只由方向相似度 决定,范数漂移被消除,防止训练过程累积的尺度漂移 。
3.3 小细节:为什么可以放在 RoPE 之前
RoPE 是纯旋转,不改变模长(02 篇 2.5 呼应 01 篇的账 的隐藏好处:)。因为先归一的后归一不影响范数,Norm(Rotate(q)) 与 Rotate(Norm(q)) 在数学上等价所以 先 ③ 后 ④ 纯属实现选择(先稳定、再注入更符合直觉,数值上也更稳)。
4. 因果掩码与 padding 掩码
4.1 因果掩码
自回归的定义:位置 只能看到 的历史(01-模型骨架与前向回路)。打分矩阵 [B, H, S, S] 的元素 是 其中 的那些列(未来位置)必须对 行不可见,否则 softmax 会把权重分给未来的自己 ,训练目标就泄漏了。
4.2 triu(1):上三角填
scores[:, :, :, -seq_len:] += torch.full((seq_len, seq_len), float( -inf ),
device=scores.device).triu(1)
torch.full(...).triu(1):构造一个 矩阵,严格上三角(主对角线以上)为 1、其余为 0:
- 加进 scores 后, 处变成 。softmax 里 是精确的
- 为什么是
triu(1)而不是triu(0):triu(0)会把主对角线(,即 自己看自己 )也遮掉。自回归允许 token 关注自己(自己的位置信息在 RoPE 相位里非常关键),所以只遮严格上三角。
小细节scores[:, :, :, -seq_len:] :推理续写时(第 6 节),scores 是 [B, H, 1, 1+past]:当前行(1 行)对着 past 全部列 + 自己 。此时只需要对新增的这一个小方块()加因果掩码。-seq_len: 把掩码只加到靠右的 seq_len 列。
4.3 padding 掩码
if attention_mask is not None:
scores += (1.0 - attention_mask.unsqueeze(1).unsqueeze(2)) * -1e9
- 入参
attention_mask形状[B, S](1=有效 token,0=padding);unsqueeze(1).unsqueeze(2)展成[B,1,1,S]后广播到[B,H,S,S]; (1.0 - mask)让 padding 列为 1,其余为 0;乘-1e9后,padding 列被加一个极大的负数,权重压到 0,却不影响其他列的 softmax 归一化( 不贡献分母)。
4.4 两种掩码的分工与数据来源
| 掩码 | 来源 | 形状 | 值 | 作用 |
|---|---|---|---|---|
| 因果掩码 | 代码构造 | [S,S] 严格上三角 | 未来不可见 | |
| padding 掩码 | 调用方传入 attention_mask | [B,1,1,S] 广播 | padding 列权重 ≈ 0 |
有意思的是:MiniMind 的训练脚本不传 attention_mask(model(input_ids, labels=labels)),所以训练时只有因果掩码生效因为 Pretrain/SFT 数据都是右 padding 到定长,padding token 处于序列末尾(位置更大),因果掩码已经把它们当作未来遮住了,loss 侧再由 -100 屏蔽(01-模型骨架与前向回路)。而 RL 训练(GRPO/PPO 的 full_mask、rollout_engine 的 padding_side=left)才真正传 mask那里的 padding 在左侧,因果掩码必须靠 attention_mask(08-偏好对齐:DPO与强化学习)。
5. Flash Attention 与手动回退
5.1 标准注意力为什么吃显存:O(S²) 的 打分矩阵
朴素实现:
scores = xq @ xk.transpose(-2, -1) # [B,H,S,S] ← 这一步物化
weights = softmax(scores / sqrt(d))
output = weights @ xv
问题在 scores:它是一个完整的 矩阵被写进显存(HBM)。一旦在长上下文(如 ):
仅打分矩阵一项就OOM。同时为了算 softmax,scores 的每行要读出来两次(一次求 max/sum、一次算 e^x),矩阵越写越大,HBM 往返越贵标准实现的瓶颈其实在内存带宽,不在算力。
5.2 FlashAttention 的原理:分块 + 在线 softmax
FlashAttention(Dao et al., 2022)没有减少计算量(仍是 ),它做的是IO 优化:把整个注意力重写成 分块(tiling) 算法,让每个块在 GPU 的片上 SRAM(几十 KB 的快速缓存)里完成 打分 → softmax → 乘 V 的全部步骤,只把块级结果写回 HBM打分矩阵 [B,H,S,S] 从不整体物化,显存降到 。
分块能成立的数学关键是在线 softmax(online softmax)。标准 softmax 要两遍(先求最大值 、再求和 ),这两遍要求整行 都在。在线 softmax 的解法:边扫边维护 运行中的最大值与和 ,遇到更大的 就把已累加部分按比例修正:
已算出的部分输出同样乘上 (rescaling)只需块内数据在 SRAM,块间只传递两个标量()。最终每个输出 token 拿到的是与标准 softmax 数学等价的结果,但全程没有 S² 矩阵落地。
5.3 PyTorch 内建的 scaled_dot_product_attention
torch.nn.functional.scaled_dot_product_attention(PyTorch ≥ 2.0)是一个融合入口:根据输入形状 / dtype / 硬件,自动选择
- flash attention 内核(NVIDIA Ampere+,bf16/fp16,且
S需满足内核对齐); - memory-efficient 内核(xformers 系,更通用的 fallback);
- math 回退(纯 PyTorch kernel,行为如 5.1 标准注意力为什么吃显存:O(S²) 的 打分矩阵 但由框架优化)。
5.4 回退成朴素注意力计算的四个条件
if self.flash and (seq_len > 1) and (not self.is_causal or past_key_value is None) \
and (attention_mask is None or torch.all(attention_mask == 1)):
output = F.scaled_dot_product_attention(xq, xk, xv,
dropout_p=self.dropout if self.training else 0.0,
is_causal=self.is_causal)
else:
# 手动回退(127–131 行): scores → 因果/attention 掩码 → softmax → @ xv
| 条件 | 含义 | 为什么需要它 |
|---|---|---|
self.flash | 构建时判定:hasattr(F,'scaled_dot_product_attention') 且 config.flash_attn | PyTorch≥2.0 版本才提供了flash attention的实现 |
seq_len > 1 | 不是单 token 续写 | 续写时打分矩阵是 [B,H,1,P+1](非方阵),flash 的 causal 内核按 方阵假设优化,非方阵不划算;且单行打分的内核启动开销反而大于手动计算 |
not is_causal or past_key_value is None | 等价于 past_key_value is None(is_causal 恒 True) | 有缓存时不能走 sdpa 的 causal 模式:causal 的意思是 行 遮列 ;而推理时当前行只该遮 未来 ,past 列必须全可见sdpa(is_causal=True) 没有未来截止点 + 过去全开这种混合掩码的表达,硬用会遮错列。这是回退最核心的原因 |
attention_mask is None or torch.all(attention_mask == 1) | 无 padding / 掩码全 1 | 见 4.3 padding 掩码:padding 掩码是 -1e9 加法;flash 内核的融合 causal 路径对非平凡 mask 支持受限,MiniMind 的策略是 有 padding 就退回手动路径自己加 |
- Flash 路径 = 训练 / 预填充(每步
seq_len>1、无 KV 缓存、无 padding 的 方阵 + 全可见 场景); - 手动路径 = 推理续写(
seq_len=1、有缓存、非方阵)或带 padding 的批处理。
故 [[#4.2 triu(1):上三角填 ]] 那个 -seq_len: 切片之所以需要是因为续写的非方阵打分只有手动路径在算。
6. KV Cache 与推理形态
6.1 推理的重复计算问题
生成(decode)是自回归循环:每一步只产生一个新 token,但朴素实现要把 prompt + 已生成全部 重新算一遍前向每步都是 ,总开销 。
关键观察在注意力的结构里:位置 的 K、V 是 、v_n=W_v x_n$$x_n 一旦生成就不变,所以历史位置的 K、V 永远不变、每步重算纯属浪费。变化的只有:每步新增一个 (新位置要问)、把这个新位置的 K、V 算出来。
KV Cache 的策略:把历史位置的 K、V 存下来,每步只算新 token 的 K、V,拼到缓存尾,打分时 新 q × 全部 K :
6.2 KV Cache实现代码
if past_key_value is not None: # ⑤ 推理续写才进来(训练 None)
xk = torch.cat([past_key_value[0], xk], dim=1) # [B, P, 4, 96] + [B, 1, 4, 96] → [B, P+1, 4, 96]
xv = torch.cat([past_key_value[1], xv], dim=1)
past_kv = (xk, xv) if use_cache else None # 返回给上层,供下一轮输入
- 拼接在
dim=1(序列维)发生在 [[#2.3repeat_kv:4 组 KV 如何服务 8 个 Q]] 的repeat_kv之前。缓存里存的是 4 组 KV(GQA 未广播形态),每次前向只广播一次, - 训练时全程 None:
use_cache=False(训练脚本不传),past_kv=None,presents列表全是 None缓存机制零开销地 隐形 ; - 容易看错的一条:只有 K/V 被拼接,xq 从不拼接。续写时
generate每步只喂input_ids[:, past_len:](§6.5),所以 xq 始终只有新 token 那一行[B,1,8,96];点积[B,8,1,96]×[B,8,96,1+P]就是 新 q × 全部 K ,没有一行历史 q 被重算。
6.3 使用KV Cache的缓存收益
每层、每个 token 的缓存量(bf16):
| 序列长度 | KV Cache(GQA 8:4) | ×2 若用 MHA(8:8) |
|---|---|---|
| 1k | 12 MB | 24 MB |
| 8k | 98 MB | 196 MB |
| 32k | 392 MB | 785 MB |
6.4 非方阵打分的正确性
续写时打分矩阵是 [B, H, 1, P+1]1 行 × (P+1) 列,非方阵:
- 当前行(新 q)允许看所有 past 列
- 未来遮列 在续写时已经不需要了:新 token 本身就是序列最末尾
- 这也正是 5.4 回退成朴素注意力计算的四个条件 判据 有 past 就回退手动路径 的深层原因非方阵 + 只看过去 的语义,flash 内核的方阵 causal 模式表达不了。
- 新 token 的 RoPE 角 = 真实的绝对位置 (
start_pos=P时的切片),past 的 K 携带生成时的位置角。
6.5 KV Cache下的生成循环
伪代码版本的推理循环
for _ in range(max_new_tokens):
outputs = forward(input_ids[:, past_len:], past_key_values) # 每步只算新 token
next_id = sample(outputs.logits[:, -1, :]) # 采样(温度/top-p,§8 待补)
input_ids += next_id
past_key_values = outputs.past_key_values # 缓存接力
算力从 降到 ( 为总长、 为历史长):每步的打分是新 (1 行)对着缓存( 列),是矩阵-向量乘法
7 o_proj:将8 个头拼接回768维
# 手动路径尾部(131 行):softmax 权重 × 加权混合 V
output = self.attn_dropout(F.softmax(scores.float(), dim=-1).type_as(xq)) @ xv
output = output.transpose(1, 2).reshape(bsz, seq_len, -1) # [B,8,S,96] → [B,S,768]
output = self.resid_dropout(self.o_proj(output)) # o_proj: 768 → 768 方阵
attn_dropout(…).@ xv:softmax 权重与 xv 的加权混合,是手动路径里唯一 V参与计算的地方;训练时 dropout 在 softmax 后按权重随机置零,flash 路径由sdpa的dropout_p在块内等价完成([[#5.3 PyTorch 内建的scaled_dot_product_attention]] 第二点);transpose(1, 2)+reshape:把 8 个头排回一条 768 维向量(回头维,§0 的 ⑨);o_proj(768→768 方阵):把 8 个头的信息线性混合成一条残差流向量
评论