本文定位上一篇从 Token IDs 到训练 Loss把
Attention暂时看成了一个保持形状不变的黑盒:本文打开这个黑盒,集中阅读
model/model_minimind.py中的Attention.forward、apply_rotary_pos_emb和repeat_kv。目标仍然是建立同一套对应关系:代码语句 ↔ 数学运算 ↔ 张量形状组件的理论作用现代 LLM 组件
1. Attention 前向传播是一条可追踪的张量流水线
Attention.forward 接收某个 Decoder Block 中经过 RMSNorm 的隐藏状态:
表示本次调用送入模型的 Query 长度。训练时通常一次送入整段序列,此时 ;使用 KV Cache 逐 Token 解码时,通常有 。
本文使用以下符号:
| 符号 | 代码属性 | 含义 |
|---|---|---|
bsz | Batch Size | |
seq_len | 本次输入的 Token 数,即 Query 长度 | |
| Key/Value 的序列维 | 历史缓存与本次输入合并后的长度 | |
hidden_size | 模型隐藏维度 | |
n_local_heads | Query 头数 | |
n_local_kv_heads | Key/Value 头数 | |
head_dim | 单个注意力头的维度 | |
n_rep | 每个 K/V 头服务的 Query 头数, |
MiniMind 默认使用 、、、,所以:
下面的图先给出整条数据流。虚线缓存支路只在自回归生成时发挥作用;QK Norm 和 RoPE 都不处理 Value。
把这条流水线压缩成一组公式,就是:
后面的代码只是在显式构造这些量,同时不断改变它们的视图和形状。
2. Q/K/V 投影把表示空间拆成多个注意力头
Attention 初始化时创建四个无偏置线性层:
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)nn.Linear(in_features, out_features) 保存的权重形状是 [out_features, in_features],前向计算采用 。因此:
forward 先完成投影proj,例如q_proj 输出的 768 维(内存里是连续的一条),再用 view 把最后一个总维度拆成“头数 × 头维度”view 不会再学习一次投影,也不会把信息分发给独立的小网络;它只是重新解释连续内存中的维度,768 个连续的数,前 96 个算头 0,接下来 96 个算头 1,以此类推。。真正决定每个头看到什么子空间的是 、 和 中不同的行, 的第 095 行产生头 0 的 query,第 96191 行产生头 1 的。:
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)对应的形状变化是:
MiniMind 默认 ,所以 Query 投影(q_proj)前后的总维度相同;但这不等于没有发生变换。它仍然用一个 参数矩阵把每个 Token 的隐藏表示映射到了新的 Query 空间。K/V 总维度只有:
这正是 GQA 比标准多头注意力节省 K/V 参数和缓存的来源。默认配置下,四个投影的主要参数量为:
如果 K/V 也各有 8 个头,标准 MHA 对应的投影参数量则是 。
标准 MHA: GQA (MiniMind):
Q: Q0 Q1 Q2 Q3 Q4 Q5 Q6 Q7 Q: Q0 Q1 | Q2 Q3 | Q4 Q5 | Q6 Q7 │ │ │ │ │ │ │ │ └─┬─┘ └─┬─┘ └─┬─┘ └─┬─┘ K: K0 K1 K2 K3 K4 K5 K6 K7 K: K0 K1 K2 K3 V: V0 V1 V2 V3 V4 V5 V6 V7 V: V0 V1 V2 V3
8 组 K/V 4 组 K/V,每组服务 2 个 Q 头代码里的 n_rep 就是这个”每组服务几个”:
self.n_rep = self.n_local_heads // self.n_local_kv_heads # 8 // 4 = 2GQA 是 MHA 和 MQA 之间的插值:
| ~ | K/V头数量 | 说明 |
|---|---|---|
| MHA | 每个 Query 头都有独立的 K/V | |
| MQA | 所有 Query 头共享同一组 K/V | |
| GQA | 每组 K/V 服务多组 Query 头 |
3. QK Norm 与 RoPE 分别控制数值尺度和位置信息
完成多头拆分(view)后,MiniMind 先归一化 Q/K,再施加 RoPE:
xq, xk = self.q_norm(xq), self.k_norm(xk)
cos, sin = position_embeddingsxq, xk = apply_rotary_pos_emb(xq, xk, cos, sin)这里的 q_norm 和 k_norm 都是 RMSNorm(self.head_dim)。以单个 Query 头向量 为例:
代码在最后一维上计算均方根,因此每个 Token 的每个头独立计算自己的归一化尺度。与此同时, 会广播到 、 和 三个维度,即所有 Query 头共享同一组可训练缩放参数;Key 使用另一组独立的 。Value 不参与 QK Norm,因为注意力分数由 Q 与 K 的点积决定,而 V 是被权重汇总的内容。
自定义 RMSNorm 还会先把输入转成 float32 计算,再转回原来的数据类型:
def _norm(self, x): return x * torch.rsqrt( x.pow(2).mean(-1, keepdim=True) + self.eps )
def forward(self, x): return (self.weight * self._norm(x.float())).type_as(x)这一步不改变形状:
随后 apply_rotary_pos_emb 将位置旋转施加到 Q/K:
def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1): def rotate_half(x): return torch.cat( (-x[..., x.shape[-1] // 2:], x[..., :x.shape[-1] // 2]), dim=-1 )
q_embed = ( q * cos.unsqueeze(1) + rotate_half(q) * sin.unsqueeze(1) ).to(q.dtype) k_embed = ( k * cos.unsqueeze(1) + rotate_half(k) * sin.unsqueeze(1) ).to(k.dtype) return q_embed, k_embed若把一个头向量沿最后一维分成等长两半:
那么位置 的旋转可以写为逐元素运算:
cos 和 sin 的形状为 ,unsqueeze(1) 后变为 ,再通过广播同时作用于 Batch 与所有头。RoPE 不在隐藏状态上直接加一个位置向量,而是在 Q/K 空间中按位置旋转坐标,使后续点积能够表达相对位置信息。
两步的职责可以明确区分:QK Norm 控制参与点积的向量尺度,RoPE 改变参与点积的方向以编码位置。它们都只处理 Q/K,也都保持张量形状不变。
4. KV Cache 与 GQA 在不同维度上减少重复计算
RoPE 之后,代码先拼接历史缓存,再展开 GQA 的 K/V 头:
if past_key_value is not None: 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 = xq.transpose(1, 2)xk = repeat_kv(xk, self.n_rep).transpose(1, 2)xv = repeat_kv(xv, self.n_rep).transpose(1, 2)假设缓存中已有 个历史 Token,本次又输入 个 Token,那么拼接后的长度为:
缓存的形状是:
这两个缓存解决的是“时间维度上的重复”:历史 Token 的 K/V 已经算过,下一步生成时直接复用即可,不必把整个前缀重新送过模型。每层缓存的元素数为:
MiniMind 的 ,因此其 KV Cache 元素数只有同头数 MHA 的一半。
但是 Query 有 个头,缓存中的 K/V 只有 个头。点积计算前,repeat_kv 在头维度上把每个 K/V 头逻辑展开 次:
def repeat_kv(x, n_rep): 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) )默认的对应关系是:
Query heads: Q0 Q1 | Q2 Q3 | Q4 Q5 | Q6 Q7K/V heads: K0 K0 | K1 K1 | K2 K2 | K3 K3group: └─ G0 ─┘ └─ G1 ─┘ └─ G2 ─┘ └─ G3 ─┘用形状表示为:
Query 不需要复制,只需转置:
GQA 解决的是“头维度上的冗余”:一组 Query 头共享同一组 K/V 表示。repeat_kv 是为了让普通批量矩阵乘法能够按相同头数执行,并不表示缓存中真的保存了 份 K/V;past_kv 在展开之前就已经建立,仍然只保存 个头。
还要注意,代码默认 能被 整除,因为 n_rep 使用整数除法。若自定义配置破坏了:
展开后的头数就无法正确匹配 Query。
5. 缩放点积、掩码与 Softmax 共同得到注意力权重
经过转置与 GQA 展开后,三个张量的形状为:
手动分支直接实现了缩放点积注意力:
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)
if attention_mask is not None: scores += ( 1.0 - attention_mask.unsqueeze(1).unsqueeze(2) ) * -1e9
weights = F.softmax(scores.float(), dim=-1).type_as(xq)output = self.attn_dropout(weights) @ xv首先计算每个 Query 与全部 Key 的相似度:
若 Q/K 各维度具有近似相同的尺度,未经缩放的点积方差会随 增长。除以 可以避免 Softmax 输入随头维度变大而过度极端。
因果掩码把“当前 Query 不应该看到的未来 Key”加为 。没有缓存时,,它是熟悉的上三角矩阵:
有缓存时,scores[:, :, :, -seq_len:] 只在本次新增的 区域加入上三角掩码。所有历史列都保持可见,因为历史 Token 对当前 Query 来说都位于过去。逐 Token 解码时 ,这个 上三角矩阵为 0,当前 Token 可以关注全部 个 Key。
外部 attention_mask 通常形如 ,两次 unsqueeze 后广播为 。其中 0 对应的 Key 位置被加上一个极小值,从而在 Softmax 后取得近似 0 的权重。这里的核心不是“用 0 乘掉分数”,而是在 Softmax 之前对不合法位置施加加性掩码:
最后沿 Key 维归一化,并对 Value 加权求和:
scores.float() 让 Softmax 在 float32 中计算,降低半精度指数运算溢出或下溢的风险;随后 type_as(xq) 再把权重转回 Q 的数据类型。attn_dropout 只在训练模式下随机丢弃部分注意力权重,不改变形状。
预训练脚本调用模型时没有显式传入 attention_mask。PretrainDataset 采用右侧 Padding,真实 Token 无法越过因果掩码看到未来的 PAD;PAD 位置虽然可以看到前文,但其标签被设为 -100,不会进入语言模型 Loss。因此这一路径仍然成立。若 Batch 使用左侧 Padding 或更复杂的有效区间,就应正确传入 attention_mask。
6. Flash Attention 是同一数学运算的融合实现
当运行环境支持且配置启用 Flash Attention 时,MiniMind 会优先调用 PyTorch 的融合算子:
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: # 手动计算 scores、mask、softmax 和 output ...这不是另一种注意力目标。两条分支都在计算:
差别在工程实现:融合算子可以分块完成点积、Mask、Softmax 与 Value 汇总,避免把完整的 分数矩阵长期写入显存。序列越长,这种中间张量的显存开销越明显。
MiniMind 的条件也揭示了两条分支各自常见的使用场景:
| 场景 | 常见分支 | 原因 |
|---|---|---|
| 整段预训练,未传 Padding Mask | Flash 分支 | seq_len > 1、无历史缓存且所有位置有效 |
| 带非全 1 Padding Mask 的输入 | 手动分支 | 当前条件不把该 Mask 交给融合算子 |
| 已有 KV Cache 的增量生成 | 手动分支 | past_key_value is not None |
| 单 Token 解码 | 手动分支 | seq_len == 1 |
所以阅读手动分支最容易看懂数学过程,但不应据此认为训练一定真的显式保存了 scores;实际是否融合取决于配置、运行环境与输入条件。
7. 合并多头后才回到 Decoder Block 的残差主干
每个 Query 头都已得到自己的上下文向量后,代码把头维移回末尾并合并:
output = output.transpose(1, 2).reshape( bsz, seq_len, -1)output = self.resid_dropout(self.o_proj(output))return output, past_kv对应的形状变化为:
多头结果不是求平均,而是先拼接成 维向量,再由 混合各头的信息并映射回模型隐藏维度。由于输出仍为 ,它才能与进入 Attention 前的残差分支相加。
不过,Attention.forward 本身只返回 Attention 输出,没有在内部执行残差加法。加法位于外层 MiniMindBlock.forward:
residual = hidden_stateshidden_states, present_key_value = self.self_attn( self.input_layernorm(hidden_states), position_embeddings, past_key_value, use_cache, attention_mask)hidden_states = residual + hidden_states因此完整关系是:
resid_dropout 的名字表示它作用于即将进入残差加法的 Attention 输出,而不是说残差连接发生在 Attention 类内部。把模块边界分清,阅读代码调用链时就不容易把两个不同位置的 RMSNorm、Dropout 与 Add 混在一起。
8. 两条形状轨迹连接训练与增量解码
为了让每个维度都能实际核对,下面使用一个缩小后的配置:
整段输入 时,形状轨迹如下:
| 阶段 | Q | K | V 或输出 |
|---|---|---|---|
| 输入 | — | — | |
| 线性投影 | [2,5,64] | [2,5,32] | [2,5,32] |
| 拆分头 | [2,5,4,16] | [2,5,2,16] | [2,5,2,16] |
| QK Norm + RoPE | [2,5,4,16] | [2,5,2,16] | [2,5,2,16] |
| GQA + 转置 | [2,4,5,16] | [2,4,5,16] | [2,4,5,16] |
| 注意力分数 | [2,4,5,5] | — | — |
| 每头上下文 | — | — | [2,4,5,16] |
| 合并头与输出投影 | — | — | [2,5,64] |
| 返回缓存 | — | [2,5,2,16] | [2,5,2,16] |
注意返回缓存仍然只有 2 个 K/V 头,而不是 GQA 展开后的 4 个头。
现在假设这 5 个 Token 已经写入缓存,下一次只输入 1 个新 Token。此时 、、:
| 阶段 | 形状 |
|---|---|
| 新输入 | [2,1,64] |
| 新 Q / K / V 投影 | [2,1,64] / [2,1,32] / [2,1,32] |
| 拼接后的紧凑 K/V | [2,6,2,16] |
| GQA 展开后的 K/V | [2,4,6,16] |
| Query | [2,4,1,16] |
| 注意力分数 | [2,4,1,6] |
| Attention 输出 | [2,1,64] |
| 返回的新缓存 | K/V 各 [2,6,2,16] |
这条轨迹说明了 KV Cache 的核心收益:第 6 个 Token 的 Query 仍需与 6 个 Key 比较,但前 5 个 Token 的 K/V 不再重复投影。生成长度增长时,每一步只为新增 Token 计算新的 Q/K/V,再把 K/V 追加到缓存中。
回到源码,可以把整个 Attention.forward 读成一句话:先将 投影为多头 Q/K/V,用 QK Norm 稳定点积、用 RoPE 注入位置,再拼接紧凑的历史 K/V 并按 GQA 展开,经过带掩码的缩放点积注意力汇总上下文,最后合并各头并投影回 维。
下一篇Attention 已经从黑盒变成了一条完整的张量流。下一篇继续阅读同一个 Decoder Block 中的
FeedForward、RMSNorm 与残差连接,并进一步区分普通 SwiGLU FFN 和 MiniMind 可选的 MoE 路径。