现代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)882017 原版
MQA(Multi-Query)81PaLM(省到极致,质量有损)
GQA(Grouped-Query)84(分 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 / 硬件,自动选择

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_attnPyTorch≥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            # 返回给上层,供下一轮输入
  1. 拼接在 dim=1(序列维)发生在 [[#2.3 repeat_kv:4 组 KV 如何服务 8 个 Q]] 的 repeat_kv 之前。缓存里存的是 4 组 KV(GQA 未广播形态),每次前向只广播一次,
  2. 训练时全程 None:use_cache=False(训练脚本不传),past_kv=None,presents 列表全是 None缓存机制零开销地 隐形 ;
  3. 容易看错的一条:只有 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)
1k12 MB24 MB
8k98 MB196 MB
32k392 MB785 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 个头的信息线性混合成一条残差流向量