Transformer 与现代大模型架构 Transformer & Modern LLM Architecture
这是整条路线的重心。2026 年的架构演化几乎全部发生在这一层:注意力的计算方式(FlashAttention → 稀疏 / 线性)、参数的组织方式(稠密 → 细粒度 MoE)、位置编码的外推能力(RoPE → 各种缩放变体)、优化的数值特性(AdamW → Muon)。把这些「为什么」讲清楚,你就具备了读懂任何新模型技术报告的能力。
阶段总览
- 能手推 Scaled Dot-Product Attention 与多头注意力,说清每一步的维度与复杂度
- 能从零实现一个可训练的 MiniGPT,并解释每个模块存在的理由
- 掌握 RoPE / ALiBi 等位置编码,理解长上下文外推的原理与手段
- 理解 MHA / MQA / GQA 的取舍,能算清 KV Cache 的显存占用
- 理解 MoE 的路由、负载均衡与专家并行,能解释为什么 2026 年主流都是 MoE
- 理解 FlashAttention 为什么快(IO 感知,不是近似),以及线性注意力 / Mamba 的动机与代价
- 能读懂一份主流模型的技术报告,并做出选型判断
| 周次 | 主题 | 交付物 |
|---|---|---|
| 第 1 周 | Attention 本质与复杂度 | 手推 + numpy 实现 + 复杂度分析 |
| 第 2 周 | Transformer 全结构 | 从零实现可训练 MiniGPT |
| 第 3 周 | 现代组件与位置编码 | RoPE / RMSNorm / SwiGLU 手写 + 外推实验 |
| 第 4 周 | MoE 与高效注意力 | 阅读 DeepSeek / Qwen 技术报告并写笔记 |
| 第 5 周 | 长上下文与推理成本 | KV Cache 显存测算 + GQA 对比实验 |
| 第 6 周 | Scaling Law 与模型横评 | 选型决策文档 |
1. Attention:一切的起点
学习路径
- 读 1.1:默写 softmax(QKᵀ/√d_k)V,背下 ÷√d_k 防 softmax 饱和与 O(T²·d)
- 跑 numpy 实现验证每行权重和为 1、因果掩码让第一行只关注自己
- 实测序列长度翻倍耗时约 4–16 倍,坐实 T² 复杂度
- 对接 M8:手写 GPT 时实现并校验因果注意力的数值正确性
核心知识点详解
- 缩放 1/√d_k 是防 softmax 饱和:Q、K 各维独立方差为 1 时,点积方差约 d_k;d_k=128 时点积落在很宽的数值区间,softmax 被推入饱和区(梯度≈0)。除以 √d_k 把方差拉回 1,这就是「缩放」的全部理由。PyTorch 的
F.scaled_dot_product_attention自动处理,手写时最容易漏的就是这一步。 - 因果掩码要在 softmax 前加:把未来位置的 score 置为
-1e9(而非 -∞ 或 0),softmax 后权重≈0,注意力退位只看过去。用np.tril(np.ones((T,T)))生成下三角掩码。常见坑:把未来位置掩成 0 而不是 -1e9,softmax 仍会把权重分给未来 token。 - O(T²·d) 是长上下文的靶心:序列翻倍(256→1024)耗时约 4–16 倍,实测接近 T² 的 16 倍。FlashAttention、上下文并行、GQA 这些优化本质上都在攻击这个平方项。
- 温度 τ 控制检索锐利度:把 softmax 前分数统一除以 τ:τ 越小分布越尖锐(近似 argmax 检索),τ 越大越接近平均池化。默认 1/√d_k 就是一个固定温度;YaRN 外推里调注意力温度,本质就是调这个 τ 以补偿未见过的长距离。
学习路径
- 读 1.2:背下 QKV 合并投影 4d² 算力、注意力算力与头数无关、KV Cache 显存公式
- 跑 MultiHeadAttention 前向,验证 SDPA 自动选内核与输出 shape
- 用 attn_flops 函数验证 h=8/16/32/64 时注意力 FLOPs 不变
- 对接 M8:手写 MHA 并与 PyTorch SDPA 对齐输出,说明 QK-Norm 的作用
核心知识点详解
- QKV 合并投影 4d² 是算力大头:d_model=d、
nn.Linear(d, 3d)一次算完 Q/K/V,参数 3d·d=3d²,再加输出投影 d·d=d²,合并后为 4d²。相比 MLP 的 8d²(FFN 部分),注意力整体约占 1/3 的参数量。 - 注意力算力与头数无关:多头只是把 d 切成 h 份并行,总 FLOPs 仍等于单头:O(T²·d)。头数改变的是每头子空间与投影矩阵形状,不是总计算量。面试高频题:h=8/16/32/64 时注意力 FLOPs 不变。
- K、V 头数与 KV Cache 大小成正比:KV Cache 显存 = 2·L·kv_heads·head_dim·seq·batch·dtype_bytes(系数 2 是 K 和 V 各一份)。MQA 只有 1 个 KV 头,GQA 用 g 个 KV 头共享给 h 个 query 头,KV 显存随 g/h 下降。
- QK-Norm 防熵坍缩:对 Q/K 各自先做归一化(RMSNorm/LN)再算注意力,防止深层 logits 方差过大导致注意力熵坍缩成 one-hot。训练早期尤其关键。常见坑:在
nn.MultiheadAttention与手写实现之间对齐输出时,忘了 QK-Norm 或 dropout 开关导致数值对不上。
学习路径
- 读 1.3:区分 MHA / MQA / GQA,理解 KV Cache 正比于 KV 头数、memory-bound 推理
- 用 kv_cache_bytes 对比 MHA/GQA/MQA 在 32K 上下文下的显存(约 64/16/2 GB)
- 读 MLA 的低维潜向量 c_kv 方案,对比其与 GQA 的 KV 量级
- 对接 M8:手写 GPT 时选用 GQA,权衡推理显存与吞吐
核心知识点详解
- KV Cache 正比于 KV 头数,与 query 头数无关:KV 显存 = 2·L·kv_heads·head_dim·seq·batch·dtype。MHA 中 kv_heads=h;MQA 只留 1 个;GQA 用 g 个 KV 头共享给 h 个 query 头。故 GQA 相对 MHA 的 KV 显存 = g/h(如 h=32、g=4 则降为 1/8)。
- 推理是 memory-bound:解码每步只算一个 token 的注意力,算术强度(FLOPs/字节)很低,瓶颈在搬 KV Cache 的带宽而非计算。因此省 KV 显存(GQA/量化/分页)直接转化为吞吐与并发提升。
- 32K 上下文三种方案的量级:取 L=40、head_dim=128、bf16 时:MHA 约 64 GB,GQA(g=8) 约 16 GB,MQA 约 2 GB(约 64:16:2)。
kv_cache_bytes函数一键可比。 - MLA 用低维潜向量 c_kv 更进一步:DeepSeek 的做法:把 K、V 压缩到一个低维潜向量 c_kv 再上采样回 head_dim,KV 显存再降一个量级。常见坑:选 GQA 时组数 g 不能简单整除导致 K/V 布局错位,加载开源权重前先核对
num_key_value_heads与 head 切分顺序。
1.1 本质:带权重的信息检索
注意力可以完全脱离 Transformer 来理解:你有一个查询(Query),一堆可选的键(Key)与对应的值(Value);用 Q 与每个 K 算相似度,softmax 归一化成权重,再对 V 加权求和。就这么简单。搜索引擎、推荐召回、记忆检索本质上都是这个模式。
pythonimport numpy as np
def softmax(x, axis=-1):
x = x - x.max(axis, keepdims=True)
e = np.exp(x)
return e / e.sum(axis, keepdims=True)
def scaled_dot_product_attention(Q, K, V, mask=None):
"""
Q: (..., Tq, d_k) K: (..., Tk, d_k) V: (..., Tk, d_v)
复杂度: O(Tq * Tk * d_k) —— 序列长度是平方级,这是所有优化的靶心
"""
d_k = Q.shape[-1]
scores = Q @ K.swapaxes(-1, -2) / np.sqrt(d_k) # 缩放:防止点积随维度增大而爆炸
if mask is not None:
scores = np.where(mask, scores, -1e9) # 因果掩码:不能看未来
weights = softmax(scores, axis=-1)
return weights @ V, weights
T, d = 8, 16
Q = np.random.randn(1, T, d); K = np.random.randn(1, T, d); V = np.random.randn(1, T, d)
causal = np.tril(np.ones((T, T), dtype=bool))
out, w = scaled_dot_product_attention(Q, K, V, causal)
print("output:", out.shape, " attention weights:", w.shape)
print("每行权重和 =", w[0].sum(-1).round(4)) # 应为全 1
print("第一行只关注自己:", w[0, 0].round(3)) # 因果掩码的效果
把 softmax 前的分数统一除以温度 τ,等价于缩放 logits:τ 越小分布越尖锐(近似 argmax 检索),τ 越大越平滑(近似对 V 做平均池化)。注意力默认的 1/√d_k 就是一个固定温度;长上下文里常叠加「注意力温度缩放」(YaRN 的做法)来补偿未见过的长距离,本质就是在调这个 τ。
python# 实测注意力的 T^2 复杂度:序列翻倍,耗时约 4 倍
import numpy as np, time
def attn_np(Q, K, V, causal):
d = Q.shape[-1]
s = Q @ K.swapaxes(-1, -2) / np.sqrt(d)
if causal:
T = s.shape[-1]
s = np.where(np.tril(np.ones((T, T), dtype=bool)), s, -1e9)
w = np.exp(s - s.max(-1, keepdims=True)); w /= w.sum(-1, keepdims=True)
return w @ V
for T in (256, 512, 1024):
Q = np.random.randn(1, T, 64).astype(np.float32)
t0 = time.perf_counter()
for _ in range(5): attn_np(Q, Q, Q, True)
print(f"T={T:5d} 平均 {1000*(time.perf_counter()-t0)/5:7.2f} ms")
# 预期:T 从 256 到 1024(4 倍),耗时约 12 到 16 倍(接近 T^2 的 16 倍)
# 结论:注意力是序列长度的二次函数,长上下文的一切优化都在攻击这个平方项
1.2 多头注意力与复杂度账本
单头注意力只能学一种「相关性模式」。多头把 d_model 切成 h 份,每个头在自己的子空间里独立做注意力,再把结果拼接投影回来——相当于并行地从多个视角检索信息。
pythonimport torch, torch.nn as nn, torch.nn.functional as F
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, n_heads=8, dropout=0.0):
super().__init__()
assert d_model % n_heads == 0
self.h = n_heads
self.d_k = d_model // n_heads
self.qkv = nn.Linear(d_model, 3 * d_model, bias=False) # 一次算完 q/k/v,更快
self.out = nn.Linear(d_model, d_model, bias=False)
self.drop = nn.Dropout(dropout)
def forward(self, x, causal=True):
B, T, C = x.shape
q, k, v = self.qkv(x).chunk(3, dim=-1)
# (B, T, C) -> (B, h, T, d_k)
q, k, v = [t.view(B, T, self.h, self.d_k).transpose(1, 2) for t in (q, k, v)]
# PyTorch 2.x 的 SDPA 会自动选择 FlashAttention / Memory-Efficient 内核
y = F.scaled_dot_product_attention(q, k, v, is_causal=causal, dropout_p=self.drop.p if self.training else 0.0)
y = y.transpose(1, 2).contiguous().view(B, T, C)
return self.out(y)
mha = MultiHeadAttention()
x = torch.randn(2, 64, 512)
print(mha(x).shape) # (2, 64, 512)
| 项目 | 复杂度 | 说明 |
|---|---|---|
| 注意力计算(训练) | O(T² · d) | 序列长度平方,长上下文的核心瓶颈 |
| 注意力计算(推理) | O(T · d) 每步 | 因为有 KV Cache,只需算当前 token |
| KV Cache 显存 | O(L · T · d_kv · 2 · 精度字节) | L=层数,T=已生成长度;这是推理显存的大头 |
| 投影层(QKV/OUT) | O(T · d²) | T 较小时反而是这部分主导 |
| FFN | O(T · d · d_ff) | 通常占参数量的 2/3,是 MoE 改造的对象 |
python# 算清 KV Cache:面试与容量规划的高频题
def kv_cache_bytes(layers, seq_len, n_kv_heads, d_head, dtype_bytes=2, batch=1):
"""n_kv_heads 是 KV 头数(GQA 后小于注意力头数)"""
return batch * layers * seq_len * n_kv_heads * d_head * 2 * dtype_bytes
# 对比:32 头 MHA vs 8 头 KV 的 GQA(d_model=4096, d_head=128)
for name, n_kv in [("MHA (32 KV heads)", 32), ("GQA (8 KV heads)", 8), ("MQA (1 KV head)", 1)]:
gb = kv_cache_bytes(32, 32768, n_kv, 128) / 1024**3
print(f"{name:22s} 32K 上下文单序列 KV Cache = {gb:5.1f} GB")
# 输出:MHA ≈ 64 GB(根本放不下) / GQA ≈ 16 GB / MQA ≈ 2 GB
# —— 这就是 2026 年几乎所有模型都用 GQA 的原因:显存与质量的折中
多头注意力的算力账本:QKV 合并投影为 3·d²,输出投影 d²,合计 4·d²;注意力分数矩阵为 (B, h, T, T),显存与算力都是 O(T²·h)。用 head_dim 表示时 h·d_head = d_model,于是 QKᵀ 的算力为 2·B·T²·d_model,只由 d_model 与 T 决定,与头数无关——这正是很多人算错的地方。
python# 把「多头」算力账本写成函数:验证总注意力 FLOPs 与头数无关
def attn_flops(B, h, T, d_head):
# QK^T 与 softmaxV 各 2*B*h*T*T*d_head
return 4 * B * h * T * T * d_head
for h in (8, 16, 32, 64):
f = attn_flops(1, h, 4096, 4096 // h) # d_model 固定为 4096
print(f"h={h:3d} d_head={4096//h:5d} 注意力 FLOPs={f/1e9:8.1f} GFLOPs")
# 预期:四行数值完全相同(约 1.10e4 GFLOPs)——注意力总算力只由 d_model 和 T 决定
# 但 KV Cache 显存正比于 KV 头数:头数越多显存越大,却不增加注意力算力
另一处易被忽略:多头并没有增加算力,只增加了「表达的子空间数」。把 d_model 拆成 h 个子空间并行做注意力再拼接投影,若 h=1 则退化为单头,算力相同但表达受限。2026 的 QK-Norm(DeepSeek、Qwen 采用)在算注意力前对 Q、K 各做一次 RMSNorm,抑制 logits 过大导致的熵坍缩,是长训练稳定的关键技巧。
1.3 MHA / MQA / GQA:KV 头数如何决定显存与吞吐
MHA(Multi-Head Attention)中 KV 头数 = 注意力头数 h,每个头维护自己的 K、V。MQA(Multi-Query)把所有查询头共享唯一一组 K、V(KV 头数 = 1)。GQA(Grouped-Query)是折中:g 个 KV 头,每 h/g 个查询头共享一个 KV 头(1 < g < h)。因为 KV Cache 大小正比于 KV 头数,GQA 把 KV Cache 降到 MHA 的 g/h(如 8/64 = 1/8),而质量损失远小于 MQA。注意力分数矩阵形状固定为 (B, h, T_q, T_k),softmax 沿最后一维(key 维)进行;GQA 只是把 K/V 先 repeat_interleave 到 h 个头再算,分数矩阵形状不变。
python# GQA 的显存 / 吞吐数学(自包含)
def kv_bytes(layers, seq, kv_heads, head_dim, batch=1, dt=2):
return 2 * layers * kv_heads * head_dim * seq * batch * dt
# 公开配置整理(约数,用于量级直觉,非精确规格)
configs = {
"Llama-4-Scout (GQA)": dict(L=48, h=64, g=8, dh=128),
"Qwen3-235B-A22B (GQA)": dict(L=94, h=64, g=8, dh=128),
"DeepSeek-V3 (MLA)": dict(L=61, h=128, g=1, dh=512), # MLA 另类:KV 被压到极低维
}
for name, c in configs.items():
mha = kv_bytes(c['L'], 32768, c['h'], c['dh']) / 1024**3
gqa = kv_bytes(c['L'], 32768, c['g'], c['dh']) / 1024**3
print(f"{name:24s} 32K KV: MHA={mha:6.1f}GB 实际GQA/MLA={gqa:6.1f}GB")
# Qwen3 全家族统一 8 个 KV 头,是 2026 性价比部署的标杆配置
| 方案 | KV 头数 | KV Cache | 质量 | 吞吐 | 典型用途 |
|---|---|---|---|---|---|
| MHA | = h | 最大 | 最佳 | 最低 | 小模型 / 质量优先 |
| GQA | g (1| 中等(h/g 倍缩小) | 接近 MHA | 高 | 2026 主流大模型默认 | |
| MQA | 1 | 最小 | 略损 | 最高 | 超低延迟 / 流式场景 |
为什么 2026 年主流模型几乎全用 GQA?根因是推理是显存带宽受限(memory-bound)而非算力受限:每解码一个 token,都要把整段 KV Cache 从 HBM 读进计算单元。Cache 越小 → 单卡能塞的 batch 越大、访存越少 → 吞吐越高、延迟越低。GQA 在几乎不掉点的前提下把 Cache 砍掉 4–8 倍,是「更长上下文 + 更高并发」与「质量」之间的最优折中。MQA 更激进但质量损失更明显,只在极致的延迟敏感场景(如部分语音/流式)出现。
MLA(Multi-head Latent Attention)是 DeepSeek 的第三条路:把 K、V 投影到一个低维「潜向量」c_kv(如 d_c=512),推理时只缓存 c_kv 与解耦的 RoPE 分量,再在每步还原 K/V。因此其 KV Cache 远小于同规模 GQA,质量却接近 MHA。代价是每次解码多一次上投影计算,实现更复杂。
python# MLA vs GQA vs MHA 的 KV Cache 量级对比
def kv_gb(L, T, dims_per_token, dt=2, B=1):
# dims_per_token: 每层每 token 实际缓存的元素数
return 2 * L * T * dims_per_token * dt * B / 1024**3
L, T = 61, 32768
mha = kv_gb(L, T, 128 * 128) # 128 KV 头 x head_dim 128
gqa = kv_gb(L, T, 8 * 128) # GQA 8 个 KV 头
mla = kv_gb(L, T, 512 + 64) # 潜向量 512 + RoPE 分量 64
print(f"32K KV: MHA={mha:7.1f}GB GQA={gqa:6.1f}GB MLA={mla:6.1f}GB")
# 预期:MHA 比 GQA 高约 16 倍;MLA 与 GQA 同量级甚至更低
# DeepSeek-V3 报告其 MLA 的 KV Cache 约为 MHA 的 1/7 到 1/9
| 方案 | KV 缓存对象 | 相对 MHA | 代表 |
|---|---|---|---|
| MHA | 完整 K,V | 1× | 早期 LLaMA |
| GQA | g 组 K,V | g/h | Llama 4 / Qwen3 / Mistral |
| MQA | 单组 K,V | 1/h | 部分流式 / 语音模型 |
| MLA | 低维潜向量 c_kv | 远小于 GQA | DeepSeek-V3/V4 / Kimi |
- 默写公式:写出 Scaled Dot-Product Attention,标注 Q/K/V 形状与复杂度。判据:softmax(QKᵀ/√d_k)·V,复杂度 O(T²·d)。
- 解释缩放:为什么除以 √d_k 而不是 d_k?判据:点积方差≈d_k,除 √d_k 把方差拉回 1,避免 softmax 饱和、梯度消失。
- 改代码:把 1.1 的 numpy 实现去掉因果掩码,观察第一行权重。判据:无掩码时第一行不再集中在自己、权重摊平。
- 手算 KV Cache:80 层、GQA 8 头、head_dim 128、FP16、32K、batch=8。参考答案:2×80×8×128×32768×8×2 B ≈ 81.9 GB。
- 选型判断:要支持 1M 上下文且高并发,KV 头数往哪调?判据:压低(GQA/MQA/MLA),因为 KV Cache 正比于 KV 头数。
2. 现代架构改进(2026 视角)
学习路径
- 读 2.1:用「对照清单」读透一个现代 Decoder 块(Pre-LN/RMSNorm/SwiGLU/RoPE/GQA/FlashAttention)
- 搭建一个现代 Decoder 块并将各位置与 2017 原版逐项对比
- 对接 M8:手写 GPT 时按此清单组织 block,逐层对齐结构与参数量
核心知识点详解
- 现代块 vs 2017 原版的对照清单:Post-LN→Pre-LN、LayerNorm→RMSNorm、ReLU→SwiGLU、绝对位置编码→RoPE、MHA→GQA、完整注意力→FlashAttention/SDPA。逐项对过去才真正读懂一份 2026 技术报告。
- 块参数 ≈ 12d² 的构成:注意力投影 4d²(QKV 3d² + out d²)+ FFN SwiGLU 8d²(两个 8/3d² 门控分支再加输出 8/3d²)。归并后每块约 2·(4d²) 量级对应 12d²。叠完 embedding 与层数即可估算总参。
- RMSNorm 只除均方根、无偏置:
x / sqrt(mean(x²)+eps)·γ,相比 LayerNorm 少了下减均值与平移偏置,计算更省、还能吸收到后续运算中。现代块普遍用它做 Pre-LN 的正则位置。 - 每一行改动都有动机对应一个坑:Pre-LN 解决深层梯度爆炸、RoPE 引入相对位置、SwiGLU 提升表达且省参(hidden=8/3d 与 4d MLP 参数量等价)。常见坑:对照开源权重时,RMSNorm 的
eps(如 1e-5 vs 1e-6)和 RoPE 的base(如 10000 vs 500000)不一致,会让 logits 差出一个量级。
学习路径
- 读 2.2:理解 Top-K 路由、共享专家、负载均衡与激活比 3–4% 的量级
- 跑一个 MoE 块计算参数量与激活参数,验证"总参大但激活少"
- 对接 M8:在 architecture.md 里用 FFN 改造的视角说明 MoE 与标准 FFN 的参数量差
核心知识点详解
- Top-K 路由 + 共享专家:每个 token 经 router 打分后走 Top-K 个专家(常见 K=8),另配一个共享专家承载通用计算。激活参数 = 各 token 实际用到的专家参数量,远小于总参。
- 激活比 3–4% 的量级:总参数 671B、激活参数约 37B(DeepSeek-V3 级),激活比≈5%;更常见的「稀疏 16 专家、Top-2」会落在 3–4%。公式:激活参 = 路由专家参 × K/N + 共享参 + 注意力 + 投影。
- 负载均衡损失 L_aux=α·Σf_i·P_i:f_i 是分配给专家 i 的 token 比例,P_i 是 router 对专家 i 的平均概率;最小化 f_i·P_i 让专家负载均匀,α 通常 0.01 量级。DeepSeek 改用无辅助损失偏置(动态调整 gate 偏置)+ router z-loss 抑制数值爆炸。
- 专家容量因子 1.25–2.0 与 EP 通信:容量因子=实际容量/理论平均,>1 时多余 token 走残差 skip;因子越低越省显存但丢 token 越多。专家并行 EP 靠 all-to-all 把 token 送到对应专家所在卡,通信随专家数上升。常见坑:只算「总参大」就以为省显存,忘了激活参数与各专家权重都要驻留显存。
学习路径
- 读 2.3:理解 RoPE 相对位置旋转、PI / NTK / YaRN 外推,并知道有效窗口 ≠ 标称窗口
- 实现 RoPE 前向并测其旋转不变量,理解相对位置性质
- 跑一次长上下文外推,记录标称与有效窗口的差距
- 对接 M8:手写 GPT 的 RoPE 并按目标 base 复现,与参考实现对齐
核心知识点详解
- RoPE 是相对位置旋转:把 Q、K 的第 i 维旋转角 θ_i = base^(-2i/d),沿维度成对做 2D 旋转,使 (q·k) 只依赖相对位移 (m−n)。
base=10000、d=128时 θ 从大到小分布,高频维负责距离精细度。 - 外推 ≠ 扩展:有效窗口常常小于标称:训练只见过 4K 的 RoPE,到 8K 时未见角度插值出现的频率混叠会让位置信号错乱,有效窗口可能只有 ~5–6K。要区分「RoPE base 标称上下文」与「实际测到的有效上下文」。
- 三种外推手法的机制:位置插值 PI 把外推坐标线性压回训练范围;NTK-aware 提高 base 缩放高频;YaRN 在 NTK 基础上加温度补偿(缩放注意力温度 τ)衰减高频。三种都要配合一段长度退火微调才稳。
- ALiBi 走另一条路:不旋转,直接在线性注意力的 score 上加一个随绝对距离线性衰减的偏置项,训练时无需额外成本即能外推,但长程依赖建模不如 RoPE 灵活。常见坑:手写 GPT 加载开源权重时漏了 RoPE
base,用默认 10000 而不是模型配置的 500000+,长上下文表现直接崩。
学习路径
- 读 2.4:理解 FlashAttention 的 IO 感知分块与在线 softmax,显存 O(T²)→O(T)
- 对比朴素注意力与 SDPA 的显存/速度,测量长序列下的差异
- 完成自测:说明 RA1–FA4 迭代与滑动窗口/线性注意力的取舍
- 对接 M8:在 architecture.md 说明注意力实现为何用 SDPA/FlashAttention
核心知识点详解
- IO 感知分块:省的是显存搬运不是 FLOPs:朴素注意力要物化整块 T×T 的中间矩阵(显存 O(T²)),FlashAttention 把它切成可放进 SRAM 的块,用在线 softmax 在块间增量合并运行最大值与指数和,中间矩阵不落全局内存,显存降到 O(T)。浅显误解:它并没有减少矩阵乘法 FLOPs。
- 在线 softmax 的机制:维护
running_max m_i与归一化和l_i,每处理一块用新 max 重缩放旧项的指数,等价于全量 softmax 但可流式。参考torch.nn.functional.sdpa,其会自动在 flash / memory-efficient / math 内核间挑选。 - FA1→FA4 的演进方向:FA1 只做前向 fused、FA2 合并前反向、FA3/FA4 面向 FP8/Block(Hopper)、针对不同硬件做 kernel 级调度。性能收益主要在长序列、大头的场景,短序列收益有限。
- 线性注意力 / Mamba 是另一路线:它们想在注意力中消掉 T² 项(如滑动窗口、SSM 状态空间),把复杂度降到 O(T),换来的是软化全局依赖。Mamba 没走这步时不能完全替代注意力,工程上常二者并用。常见坑:拿短序列 benchmark 炫耀 FlashAttention,收益被忽略。
2.1 一个现代 Decoder 块长什么样
把 2017 年的原始 Transformer 和 2026 年的主流模型对比,差异集中在几个固定位置。下面这张表是「读任何新模型技术报告」的对照清单。
| 位置 | 2017 原版 | 2026 主流 | 为什么改 |
|---|---|---|---|
| 归一化位置 | Post-LN | Pre-LN(归一化在子层之前) | Pre-LN 梯度更稳,不需要精细的学习率预热 |
| 归一化类型 | LayerNorm | RMSNorm | 省掉均值计算,更快,效果相当 |
| 激活函数 | ReLU | SwiGLU / GeGLU | 门控带来更好的表达力,代价是参数量增加 |
| 位置编码 | 正余弦绝对编码 | RoPE(含各种缩放变体) | 相对位置、可外推到更长上下文 |
| 注意力头 | MHA | GQA / MQA | 大幅削减 KV Cache |
| 注意力实现 | 朴素 matmul + softmax | FlashAttention / SDPA | IO 感知的分块计算,显存线性、速度更快 |
| FFN 组织 | 稠密 FFN | 细粒度 MoE(共享专家 + 路由专家) | 总参数量大但激活参数少,性价比高 |
| 优化器 | Adam | AdamW → Muon | 解耦权重衰减;Muon 做谱范数正交化提升稳定性 |
| 注意力架构 | 纯二次注意力 | 混合:滑动窗口 / 稀疏 / 线性注意力 + 少量全局层 | 把 O(T²) 压到接近 O(T) |
| 训练精度 | FP32 | BF16 / FP8(配合 QAT) | 吞吐与显存,FP8 是 2026 年大集群的常态 |
python# 现代 Decoder 块:把上面这张表组装起来
import torch, torch.nn as nn, torch.nn.functional as F
class RMSNorm(nn.Module):
def __init__(self, d, eps=1e-6):
super().__init__(); self.eps, self.w = eps, nn.Parameter(torch.ones(d))
def forward(self, x):
dt = x.dtype; x = x.float()
return (self.w * (x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)).to(dt))
class GQAttention(nn.Module):
"""分组查询注意力:Q 有 h 个头,KV 只有 g 个(g < h),节省 KV Cache。"""
def __init__(self, d=512, h=8, g=2):
super().__init__()
self.h, self.g, self.dh = h, g, d // h
self.wq = nn.Linear(d, h * self.dh, bias=False)
self.wk = nn.Linear(d, g * self.dh, bias=False)
self.wv = nn.Linear(d, g * self.dh, bias=False)
self.wo = nn.Linear(h * self.dh, d, bias=False)
def forward(self, x):
B, T, C = x.shape
q = self.wq(x).view(B, T, self.h, self.dh).transpose(1, 2)
k = self.wk(x).view(B, T, self.g, self.dh).transpose(1, 2)
v = self.wv(x).view(B, T, self.g, self.dh).transpose(1, 2)
k = k.repeat_interleave(self.h // self.g, dim=1) # 广播到与 Q 同样的头数
v = v.repeat_interleave(self.h // self.g, dim=1)
y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
return self.wo(y.transpose(1, 2).contiguous().view(B, T, C))
class Block(nn.Module):
def __init__(self, d=512, h=8, g=2, mult=2.67):
super().__init__()
self.n1, self.attn = RMSNorm(d), GQAttention(d, h, g)
self.n2 = RMSNorm(d)
hid = int(d * mult)
self.ff = nn.Sequential(nn.Linear(d, hid, bias=False), nn.SiLU(),
nn.Linear(hid, d, bias=False)) # 简化版 SwiGLU 位置
def forward(self, x):
x = x + self.attn(self.n1(x)) # Pre-LN + 残差
return x + self.ff(self.n2(x))
print(Block()(torch.randn(2, 32, 512)).shape)
把上面那张表落到「一个块里有几个矩阵」:注意力 4 个投影(4d²)、SwiGLU FFN 3 个(3·d·d_ff ≈ 8d²)、两个 RMSNorm(各 d 个权重),全部无 bias。于是现代块总参数 ≈ 12d²,其中 FFN 占约 2/3。
python# Pre-LN vs Post-LN:看残差主路的梯度能否「直通」
import torch, torch.nn as nn
class BlockLN(nn.Module):
def __init__(self, d=256, pre=True):
super().__init__(); self.pre = pre
self.ln1, self.ln2 = nn.LayerNorm(d), nn.LayerNorm(d)
self.f1 = nn.Linear(d, d); self.f2 = nn.Linear(d, d)
def forward(self, x):
if self.pre: # Pre-LN:残差主路无归一化
y = x + self.f1(self.ln1(x))
return y + self.f2(self.ln2(y))
y = self.ln1(x + self.f1(x)) # Post-LN:归一化在残差相加之后
return self.ln2(y + self.f2(y))
for pre in (True, False):
net = nn.Sequential(*[BlockLN(256, pre) for _ in range(24)])
x = torch.randn(4, 8, 256, requires_grad=True)
net(x).mean().backward()
print(f"Pre-LN={pre} 输入梯度范数={x.grad.norm().item():.4f}")
# 预期:Post-LN 的输入梯度范数明显更小(深层衰减),Pre-LN 保持接近 1
# 这就是现代模型默认 Pre-LN、而 Post-LN 必须长 warmup 的定量原因
2.2 MoE:2026 年最主流的参数组织方式
MoE 的核心交易是:用更多总参数换取更多知识容量,但每次只激活一小部分,因此计算量(FLOPs)几乎不变。2026 年前沿模型基本全是 MoE(DeepSeek-V4 约 1.6T 总参 / 49B 激活,Qwen3.5 约 397B / 17B 激活,Kimi K3 达 2.8T 级别)。
- 路由:每个 token 交给一个小的 router 网络打分,选 Top-K 个专家(常见 K=1 共享 + 6~8 路由)。
- 负载均衡:必须加辅助损失或采用无辅助损失的偏置调整,否则少数专家吃掉所有 token,其余专家训不动。
- 专家并行(EP):专家分布在不同 GPU 上,token 需要跨卡 dispatch 与 combine,通信是主要瓶颈。
- 共享专家:设若干「总是激活」的专家承载通用知识,路由专家承载专门知识(DeepSeek 的做法)。
- 细粒度专家 + 共享专家:把专家切得更小更多,组合空间更大,是目前效果最好的配置。
python# 极简 MoE 前向:理解路由与负载均衡损失
import torch, torch.nn as nn, torch.nn.functional as F
class MoE(nn.Module):
def __init__(self, d=512, n_exp=8, top_k=2, hid=1024):
super().__init__()
self.top_k = top_k
self.router = nn.Linear(d, n_exp, bias=False)
self.experts = nn.ModuleList([
nn.Sequential(nn.Linear(d, hid, bias=False), nn.SiLU(),
nn.Linear(hid, d, bias=False)) for _ in range(n_exp)])
def forward(self, x):
B, T, C = x.shape
flat = x.view(-1, C)
logits = self.router(flat) # (N, n_exp)
probs = F.softmax(logits, dim=-1)
w, idx = probs.topk(self.top_k, dim=-1) # 每个 token 选 top_k 个专家
w = w / w.sum(-1, keepdim=True) # 权重归一化
out = torch.zeros_like(flat)
for e, expert in enumerate(self.experts):
mask = (idx == e) # 哪些 token 路由到了本专家
if mask.any():
tok = flat[mask.any(-1)]
coef = w[mask].unsqueeze(-1)
out[mask.any(-1)] += expert(tok) * coef
# 负载均衡损失:鼓励专家被均匀使用(否则会“赢者通吃”)
frac = F.one_hot(idx, self.experts.__len__()).float().mean(0).mean(0) # 实际分配比例
mean_p = probs.mean(0) # 平均路由概率
aux = (frac * mean_p).sum() * self.experts.__len__()
return out.view(B, T, C), aux
负载均衡的数学形式:设 f_i 是第 i 个专家实际分到的 token 比例,P_i 是 router 平均给第 i 个专家的概率,则辅助损失 L_aux = α · Σ_i f_i · P_i。当分配完全均匀时 f_i = P_i = 1/E,L_aux = α/E 最小;若「赢者通吃」则少数 f_i→1 而多数→0,L_aux 上升。系数 α 常取 10⁻² 量级。DeepSeek 进一步提出无辅助损失均衡:直接对 router logit 加可学习的逐专家偏置 b_i,每步把 b_i 朝「被低估的专家」方向推,从而在不引入额外损失项的前提下均衡。
python# 正确的负载均衡 aux loss(top-k 路由场景)
def aux_loss(router_logits, expert_indices, n_experts, alpha=0.01):
# router_logits: (N, E) expert_indices: (N, k)
probs = F.softmax(router_logits, dim=-1) # (N, E)
# f_i: 实际被选中的比例
one_hot = F.one_hot(expert_indices.flatten(), n_experts).float() # (N*k, E)
f = one_hot.mean(0) # (E,)
p = probs.mean(0) # (E,)
return alpha * (f * p).sum()
# 专家容量与 token 丢弃
def dispatch(tokens, idx, k, capacity_factor=1.25):
N, E = tokens.shape[0], idx.max().item() + 1
cap = int(N / E * capacity_factor * k) # 每专家上限
gate = torch.zeros(E, cap, dtype=torch.long) # 越界 token 被丢弃
cnt = [0]*E
for t, e in zip(range(N), idx[:, 0].tolist()):
if cnt[e] < cap:
gate[e, cnt[e]] = t; cnt[e] += 1
# else: dropped(容量外 token 不参与本步计算)
return gate, cnt
除负载均衡外,2026 主流 MoE 还用两类 stabilizer:① router z-loss L_z = (log Σ_e exp(logit_e))²,惩罚过大的 router logit,抑制路由熵过早坍缩,对训练稳定性帮助明显(DeepSeek / ST-MoE 采用);② 专家容量上限:每个专家只处理前 capacity = (tokens/E)·capacity_factor 个 token,超出的被丢弃或溢出到下一优先专家,保证计算量可控(代价是可能丢信息,需配合数据重洗)。
推理与服务端的 MoE 是专家并行(EP):专家分布在不同 GPU,每个 token 要先 all-to-all 把输入「派发(dispatch)」到对应专家所在卡,算完再 all-to-all「合并(combine)」回来。这一步的集合通信是 MoE 推理的主要瓶颈,且通信量随专家数、序列长度增长。因此 MoE 的部署对 NVLink / 高速互联依赖极强。
| 维度 | 稠密 Dense | MoE(如 8 专家激活 2) |
|---|---|---|
| 总参数 | N | 数倍于 N(容量大) |
| 每 token 激活参数 | N | ≈ N/4(省计算) |
| 每 token FLOPs | 高 | 低(约 1/4) |
| 显存占用 | N | ≥ N(全部专家常驻) |
| 训练通信 | DP/TP/PP | 额外 EP all-to-all |
| 质量/美元 | 中等 | 更优(同成本更高容量) |
把「省计算不省显存」量化:总参数 P_total、每 token 激活 P_act,则每 token 训练 FLOPs ∝ 6·P_act,而显存必须驻留 P_total(含所有专家)。激活比 P_act / P_total 越低越省算力,但过低会让专家欠训练。
python# MoE 经济学:激活比与「等效稠密」算力
def moe_econ(total_b, active_b, dense_equiv_b):
print(f"总参 {total_b}B / 激活 {active_b}B 激活比={active_b/total_b:.2%}"
f" 推理FLOPs≈稠密{dense_equiv_b}B的 {active_b/dense_equiv_b:.2f}x")
moe_econ(1600, 49, 400) # DeepSeek-V4 量级
moe_econ(397, 17, 70) # Qwen3.5 量级
moe_econ(2800, 60, 600) # Kimi K3 量级
# 预期:激活比约 3%-4%;用约 1/8 的稠密算力换到接近稠密大模型的质量
| 模型族 | 总参 | 激活 | KV 方案 | 备注 |
|---|---|---|---|---|
| DeepSeek-V4 | ~1.6T | ~49B | MLA | 细粒度 + 共享专家,无辅助损失均衡 |
| Qwen3.5 | ~397B | ~17B | GQA 8 头 | 多语言 + 门控线性注意力混合 |
| Llama 4 | Maverick/Scout | 部分激活 | GQA | 原生多模态,开源权重 |
| Kimi K3 | ~2.8T | ~60B | MLA/KDA | 超长上下文 + Agentic |
| GLM-5.x | MoE | — | GQA/DSA | 中文与长程 Agent |
2.3 位置编码:RoPE 与长上下文外推
注意力本身是置换不变的——打乱 token 顺序结果不变。所以必须注入位置信息。RoPE 的做法是「按位置对 Q/K 做旋转」,使得两个位置的注意力分数只依赖它们的相对距离,天然适合外推。
python# RoPE 的核心:把 (x_2i, x_2i+1) 看作复数,乘以 e^{i·m·θ}
import torch
def precompute_freqs(d_head, max_seq, base=10000.0):
theta = 1.0 / (base ** (torch.arange(0, d_head, 2).float() / d_head))
pos = torch.arange(max_seq).float()
freqs = torch.outer(pos, theta) # (T, d_head/2)
return torch.polar(torch.ones_like(freqs), freqs) # e^{i m θ}
def apply_rope(x, freqs):
"""x: (B, h, T, d_head) -> 复数形式后逐位置旋转"""
xc = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
xr = torch.view_as_real(xc * freqs[None, None, :x.shape[-2], :])
return xr.flatten(-2).type_as(x)
# 长上下文外推的三类主流手段:
# 1) 位置插值 / YaRN / NTK-aware:把训练时的位置范围“压缩”到推理范围
# 2) 注意力温度缩放:对 logits 做缩放以补偿未见过的长距离
# 3) 直接继续训练更长上下文(最可靠,也最贵)
print("RoPE 的价值在于:相似度只取决于相对距离,因此具备外推的基础")
| 长上下文手段 | 代表模型 / 论文 | 代价 |
|---|---|---|
| 位置插值(PI) | Positions Interpolation | 简单有效,但过长时性能下降 |
| NTK-aware 缩放 | NTK-Aware Scaled RoPE | 改动小,适合推理端临时扩展 |
| YaRN | YaRN: Efficient Context Window Extension | 综合效果好,需少量续训 |
| 稀疏 / 滑动窗口注意力 | Mistral 滑动窗口、Longformer | 近似全局,远距离依赖可能丢失 |
| 线性 / 混合注意力 | Gated DeltaNet、KDA、CSA/HCA | 吞吐大幅提升,表达能力需验证 |
| 完整注意力的大规模长上下文训练 | Gemini 3 / Claude 1M 上下文 | 算力成本极高,但效果最可靠 |
要把 RoPE 讲透,关键是它的旋转矩阵与相对位置性质。把 d_model 维向量按相邻两维分成 d_model/2 个二维对,第 i 对的位置 m 旋转角 θ_i,旋转矩阵是这些 2×2 块的对角拼接:R(m) = diag_i [[cos(mθ_i), −sin(mθ_i)], [sin(mθ_i), cos(mθ_i)]],其中 θ_i = base^(−2i/d_model)。对 query 与 key 分别左乘 R(m)、R(n) 后做点积,q_m̃·k_ñ = q_mᵀ·R(m)ᵀ·R(n)·k_n = q_mᵀ·R(n−m)·k_n——只依赖相对距离 (n−m),与绝对位置无关。这正是 RoPE 能外推的本质。
python# 验证 RoPE 的「相对位置性质」:q_m^T k_n 只取决于 (n-m)
import numpy as np
d = 8; base = 10000.0
theta = base ** (-np.arange(0, d, 2) / d) # (d/2,)
def rot(x, m):
# x: (d,) ; 逐对旋转
c, s = np.cos(m*theta), np.sin(m*theta)
out = x.copy()
out[0::2] = x[0::2]*c - x[2::2]*s
out[2::2] = x[0::2]*s + x[2::2]*c
return out
q, k = np.random.randn(d), np.random.randn(d)
m, n = 3, 7
dot_abs = rot(q, m) @ rot(k, n)
dot_rel = q @ rot(k, n - m) # 应完全相等
print("|q_m^T k_n - q^T R(n-m) k| =", abs(dot_abs - dot_rel).round(10))
print("=> 点积只依赖相对距离 (n-m) =", n - m)
ALiBi(Attention with Linear Biases) 走另一条路:不加任何位置嵌入,而是给注意力分数直接加一个与距离成线性、随头指数衰减的偏置 b·|i−j|(b 为每个头不同的斜率)。它不引入新参数、不依赖位置编码的学习,且天然能外推到比训练更长的上下文——但表达力弱于 RoPE,2026 年已较少作为唯一方案。
| 位置编码 | 机制 | 外推性 | 代表 |
|---|---|---|---|
| 绝对位置嵌入 | 给每个位置一个可学习向量相加 | 差,超出训练长度即失效 | 原版 Transformer |
| 正弦绝对编码 | 固定 sin/cos 注入 | 一般 | 原版、部分语音模型 |
| RoPE | 对 Q/K 做旋转,相对位置守恒 | 好(配合缩放) | LLaMA / Qwen / DeepSeek |
| ALiBi | 注意力分数加距离线性偏置 | 很好(无需缩放) | BLOOM / 早期 Falcon |
| 相对位置(T5 / Shaw) | 在注意力加相对位置表征 | 中等 | T5 / DeBERTa |
RoPE 的外推靠「缩放位置」而非改模型权重。四类主流手段:① 位置插值 PI:推理时把位置 m 换成 m/s(s>1 为外推因子),把训练区间「压缩」进已知范围,简单但过长会损失分辨率;② NTK-aware / Dynamic NTK:不压缩位置,而是增大 base,保留高频维度、拉伸低频,推理端零改动;③ YaRN:在 NTK 基础上对注意力做温度缩放(高频维度降权),通常需要少量续训,综合效果最好;④ ABF(Attention Bias Factorization,2026) 等新方法,把位置偏置分解为可学习因子,进一步减小长上下文衰减。Qwen3 / Llama 4 等 2026 模型普遍在 RoPE 上叠加上述一种或多种方案,再配合少量长上下文续训,达到 128K–1M 上下文。
python# YaRN 的核心:给低频维度按 scale 拉伸频率 + 对注意力加温度补偿
import torch, math
def yarn_inv_freq(d_head, base=10000.0, scale=8.0):
inv = 1.0 / (base ** (torch.arange(0, d_head, 2).float() / d_head))
return inv # 实际 YaRN 会对低频分量按 scale 拉伸
def yarn_attn_temp(scale):
# 位置缩放后对 logits 乘温度,补偿长距离注意力衰减
return 0.1 * math.log(scale) + 1.0
for scale in (4, 8, 16):
print(f"scale={scale:2d} 有效窗口≈训练长度x{scale} 注意力温度={yarn_attn_temp(scale):.4f}")
# 预期:scale 从 4 到 16,温度从约 1.14 增到约 1.28
# 直觉:外推越远越要「降温」注意力,缓解未见长距离带来的分布偏移
| 模型(2026) | 原生窗口 | 扩展后 | 手段 |
|---|---|---|---|
| Qwen3 系列 | 32K | 128K | NTK / Dynamic NTK + 续训 |
| DeepSeek-V3/V4 | 128K | 128K+ | YaRN + 长样本文本续训 |
| Llama 4 | 256K | 10M 级 | 位置缩放 + 长上下文训练 |
| Gemini 3 Pro | — | 1M–2M | 档案级长训练 + 原生多模态 |
| Claude 4.5 | — | 1M | 长上下文训练 + 检索友好 |
2.4 为什么 FlashAttention 是真正的革命
关键认知:FlashAttention 不是近似算法,它算出的结果与朴素注意力完全一致。它快的原因不是少算了,而是少搬运了——GPU 显存层次中,HBM(大显存)带宽远低于 SRAM(片上缓存),而朴素实现需要把 T×T 的注意力矩阵写回 HBM 再读出来。
- IO 感知(IO-aware):把 Q/K/V 分块搬进 SRAM,在片内完成 softmax 与加权求和,从不物化完整的 T×T 矩阵。
- 在线 softmax:用递推方式维护 running max 与 running sum,因此可以分块计算 softmax 而不需要全局归一化因子。
- 显存效果:注意力显存从 O(T²) 降到 O(T),这是能训 32K / 128K 上下文的直接原因。
- 实现现状:PyTorch 2.x 的
F.scaled_dot_product_attention会自动挑选 FlashAttention 或 memory-efficient 内核;也可以直接用 flash-attn 库。 - 进一步的方向:FlashAttention-3/4 针对 Hopper / Blackwell 继续优化;线性注意力与 Mamba 则从算法层面把复杂度降到近 O(T)。
FlashAttention-2(2023) 的改进在「并行度」:把 work 在序列维与头维上更均衡地铺到 GPU 的 warp 上,消除 FA1 里 warp 间不必要的共享内存通信,并去掉 softmax 的冗余 rescaling,吞吐约比 FA1 高 2×。FlashAttention-3(2024) 专攻 Hopper(H100):用异步 Tensor Core + TMA(Tensor Memory Accelerator) 把矩阵乘与 softmax 在硬件层面重叠起来,并原生支持 FP8 Tensor Core,相比 FA2 再快约 1.5–2×。FlashAttention-4 面向 Blackwell(B200),进一步把异步与 warp-specialization 推到极致。关键不变式:三者算出的结果与朴素注意力逐位一致——它们是精确算法,不是近似。
| 版本 | 核心改进 | 相对提速 | 硬件 |
|---|---|---|---|
| FA1 | IO 感知分块 + 在线 softmax | 基线 | Ampere 及更早 |
| FA2 | 更优 warp 并行 + 去冗余 rescale | ~2× over FA1 | Ampere / Ada |
| FA3 | 异步 GEMM/softmax 重叠 + FP8 | ~1.5–2× over FA2 | Hopper (H100) |
| FA4 | 更强异步 + warp 专业化 | 更高(FP8/FP4) | Blackwell (B200) |
当序列长度再上一个量级(1M+),即便 O(T) 显存的 FlashAttention 也扛不住,于是出现稀疏 / 滑动窗口 与线性 / 状态空间 两条路线。滑动窗口注意力(Mistral 的 Sliding Window):每个 token 只 attend 前后 w 个邻居,复杂度降到 O(T·w),再配合少量全局 token 或层级窗口覆盖长距离;块稀疏(Longformer / BigBird)用固定稀疏图案兼顾局部与全局。这些是有损近似,长距依赖可能丢失。
线性注意力与 Mamba / SSM 更激进:把 softmax 换成「可分解为递推」的形式,使注意力变成一个大小为 d×d(或选择性状态空间)的循环状态,复杂度从 O(T²) 降到 O(T·d²) ≈ O(T),且显存恒定、可流式处理百万 token。代价是无法精确「检索」任意遥远的某个 token(线性注意力对历史做低秩压缩,会丢信息),在需精确长程检索的任务上弱于二次注意力。2026 年的主流取舍是混合架构:少量全局二次注意力层(保住检索能力)+ 大量线性/SSM 层(压低长序列成本),如 Gated DeltaNet、Griffin、Jamba(Mamba+Transformer),以及 Qwen3 引入的门控线性注意力变体。
为什么 FlashAttention 快,要用「HBM 访存量」而不是 FLOPs 解释:朴素实现每算一层都要把 T×T 的分数矩阵写回 HBM 再读出来;FlashAttention 只把 O(T·d) 的输入搬进 SRAM,中间矩阵留在片上。当 T 很大时访存成为瓶颈,省访存就等于省时间。
python# 朴素注意力 vs FlashAttention 的 HBM 访存(量级对比)
def naive_traffic(B, h, T, d_head, dt=2):
return 2 * B * h * T * T * dt # 分数矩阵写回 + 读回
def flash_traffic(B, h, T, d_head, dt=2):
return 3 * B * h * T * d_head * dt # 只搬 QKV 输入
for T in (4096, 32768):
n = naive_traffic(1, 32, T, 128) / 1024**3
f = flash_traffic(1, 32, T, 128) / 1024**3
print(f"T={T:6d} 朴素={n:8.2f}GB Flash≈{f:6.3f}GB 访存比={n/f:7.1f}x")
# 预期:T=4096 时访存比约 1000x,T=32768 时约 8600x
# 说明:T 越大,FlashAttention 省下的 HBM 搬运越可观——这才是它「快」的本质
2026 的混合架构把「检索」与「效率」分开:以 Gated DeltaNet / KDA / 门控线性注意力 作为主体层(O(T)、常数显存、可流式),每隔若干层插入少量全局二次注意力层保住精确检索能力。Qwen3-Next 一类模型采用「3:1 或更高的线性:全局层配比」,长序列吞吐可提升数倍,而需要精确 recall 的任务靠那几层全局注意力兜底。
- 对照表默写:说出 2017 原版到 2026 主流在「归一化位置/类型、激活、位置编码、注意力头、实现」五处的变化及理由。判据:Pre-LN、RMSNorm、SwiGLU、RoPE、GQA、FlashAttention 各能说一句 why。
- 算激活比:某 MoE 总参 600B、激活 30B,激活比多少?参考答案:5%;每 token 训练 FLOPs 正比于激活参数量。
- MoE 失效诊断:训练中某专家吃掉 80% token 怎么排查?判据:看 f_i·P_i 与 router 熵,加 aux loss / 偏置调整 / 调 capacity_factor。
- RoPE 外推:解释 PI 与 NTK-aware 的差别。判据:PI 压缩位置(改输入),NTK 增大 base(改频率),YaRN 再叠温度缩放。
- FlashAttention 判断:它省的是 FLOPs 还是访存?判据:访存(IO);结果与朴素注意力逐位一致,不是近似。
- 混合架构取舍:什么时候不敢用线性注意力?判据:需要精确长程检索(大海捞针、跨段推理)时全局注意力更稳。
3. 主流模型架构横评与选型
学习路径
- 读 3.1:梳理 2026 六家前沿与开源阵营,记 DeepSeek MLA+MoE / Qwen MoE / Kimi 长上下文
- 对比开源 vs 闭源在编码/数学/Agentic/多模态上的差距分布
- 对接 M8:选定一个具公开权重的 GPT 系开源模型作为对齐与验证目标
核心知识点详解
- 2026 六大前沿的归纳句:Claude/GPT/Gemini:闭源护栏与 Agentic 强;DeepSeek:MLA+MoE+Muon 把单卡推理与训练成本压下来;Qwen:MoE+线性注意力兼顾性价比;Kimi:超长上下文。各记一句「架构点+适用场景」便够选型用。
- 开源逼近闭源,但差距在长尾:编码/数学/工具调用已接近,差距集中在 Agentic 稳定性、多模态与专项产品化。选型别只看榜单均值,要看你的负载分布。
- 对齐任务要选「GPT 系且权重可复现」:M8 要对齐开源权重,需满足:具公开权重、GPT(Decoder-only)结构、tokenizer/RoPE base/eps 文档齐全。DeepSeek/Qwen/LLaMA 系都合适,确认 word embedding 与 output head 是否 weight tying。
- 防打脸要点:确认所选权重支持你的 RoPE base、词表与分词器,避免「实现完才发现 base 不兼容」再返工。常见坑:选了个闭源或没有 config 元数据的权重,对齐时无从下手。
学习路径
- 读 3.2:背下选型决策清单(先定硬约束→加权打分→LiteLLM 网关→降级链→影子灰度→小规模评测)
- 给定一个业务约束,按清单选型并解释加权依据
- 对接 M8:确认所选权重支持你的 RoPE base / 词表,避免实现与选型冲突
核心知识点详解
- 先定硬约束再加权打分:硬约束如延迟上限、每千次调用成本、是否允许数据出境、是否必须自托管、合规。先做
constraint_filter排除,再对剩下的能力维度(长上下文/工具调用/结构化输出/多语言/多模态)加权打分排序,避免「很优但根本不能用」。 - 用真实数据做小规模评测,别靠排行榜:取 100–300 条自己的真实样本先跑一遍,而不是直接用公开榜单。排行榜三条失败模式:过度拟合榜单题、测试集污染、与你的任务分布不匹配。
- 算总拥有成本 TCO:API 单价 × 预估 token 量,或自托管 GPU 小时成本 + 运维人力 + 工程时间。开源省的是单价,不是总成本。
- 用网关隔离,留降级链:用 LiteLLM / 统一网关抽象 provider,配 fallback 与影子灰度,避免供应商锁定。常见坑:只算单价不算长尾与切换成本,结果被单一模型绑架。
核心知识点详解
- 日 500 万次短分类的选型口径:先算吞吐与成本上界:日 5e6 次 × 每次 ~300 token ≈ 1.5e9 token/day;月成本 = token 量 × 单价。据此倒推需要自托管还是 API,再套硬约束。
- 月成本估算公式:monthly_cost = 日调用量 × 每次 token × 每 token 单价 × 30。短路分类优先选模型延迟低、单 token 便宜的中小模型,不必上大模型。
- 三条排行榜失败模式:① 排行榜题被模型背过(污染);② 只看均值掩盖长尾;③ 与你的负载分布错配。用你自己数据做 100–300 条盲评最可靠。
3.1 2026 年的模型格局
2026 年的一个重要变化是:前沿不再是两家垄断。按公开的智能指数,已有六家实验室拿出前沿级模型,同时开源权重阵营在编码与数学上逼近闭源,差距主要残留在Agentic 工具使用与最难推理以及视频类多模态上。
| 模型族 | 代表 | 架构要点 | 典型场景 |
|---|---|---|---|
| OpenAI GPT 系列 | GPT-5 / GPT-5.x | 闭源,自动路由 + 可调推理强度 | 通用 API、复杂 Agentic 工具使用 |
| Anthropic Claude | Claude Sonnet 4.5 / Opus 4.7 | 闭源,长上下文 1M,工具调用稳定性强 | 编码、长文档、MCP 生产 Agent |
| Google Gemini | Gemini 3 Pro / 3.5 | 原生多模态,2M 上下文 | 视频音频、科研、Workspace 集成 |
| DeepSeek | V3 / V4 | MLA + MoE + Muon,GRPO 后训练 | 自托管前沿、成本敏感场景 |
| Qwen | Qwen3 / 3.5(235B / 397B MoE) | MoE + 门控 DeltaNet 线性注意力,多语言强 | 多语言部署、Apache 许可要求 |
| Kimi | K2 / K3 | MoE + MLA/KDA,超长上下文 | 长文档、Agentic 编码 |
| GLM | GLM-4.5 / 5 / 5.2 | MoE + DSA,GRPO 起步后长程回归 Critic | 中文场景、Agent 长程任务 |
| Llama / Gemma | Llama 4 / Gemma 3-4 | 开源友好许可,端侧与自托管 | 私有化部署、微调底座 |
| MiniMax / 混元 / MiMo | M2.5 / 混元 / MiMo | 开源权重,编码与 Agentic 表现突出 | 性价比自托管、垂直微调 |
python# 把「选型」变成可复算的数:按负载加权打分
MODELS = {
# name: (能力分 0-10, 每百万 token 输入价 USD, 上下文 K)
"gpt-5": (9.6, 1.25, 400),
"claude-opus-4.7": (9.5, 3.00, 1000),
"gemini-3-pro": (9.4, 1.25, 2000),
"deepseek-v4": (8.9, 0.28, 128),
"qwen3.5-235b": (8.6, 0.20, 256),
}
def score(m, w_quality=0.6, w_cost=0.3, w_ctx=0.1, ctx_need=128):
q, price, ctx = MODELS[m]
cost = max(0.0, 1 - price / 3.0) # 3 USD/M 封顶归一化
cover = min(1.0, ctx / ctx_need) # 上下文是否够用
return w_quality * q / 10 + w_cost * cost + w_ctx * cover
for m in MODELS:
print(f"{m:18s} 综合分={score(m):.3f}")
# 预期:质量权重高时旗舰领先;把成本权重调高,开源 MoE 反超
# 结论:选型结论高度依赖权重 w_*,必须按自己的负载设权重,不能抄榜单
| 负载类型 | 最吃的能力 | 优先考虑 |
|---|---|---|
| 超长文档 / 合同审阅 | 长上下文 + 检索质量 | Gemini 3 / Claude 4.5 / Kimi K3 |
| 多步工具调用 Agent | Agentic 稳定性 + 工具调用 | GPT-5 / Claude 4.5 / GLM-5.x |
| 大批量分类抽取 | 成本 + 一致性 | DeepSeek-V4 / Qwen3.5 |
| 私有化 / 数据不出境 | 可自托管 + 许可 | Llama 4 / Qwen3.5 / Gemma |
| 端侧 / 低延迟 | 小模型 + 量化 | Qwen3 小杯 / Gemma 3-4 |
3.2 选型决策清单
python# 用统一网关把「换模型」变成改一行配置 —— 生产系统的必备抽象
# pip install litellm
from litellm import completion
MODEL_MAP = { # 按任务路由,而非全站一个模型
"router": "gpt-5-mini", # 意图识别、简单分诊:用小模型
"reasoning": "claude-opus-4-7", # 最难推理:用最强模型
"bulk": "deepseek-v4", # 大批量处理:用便宜且够用的模型
"local": "ollama/qwen3:8b", # 敏感数据:本地推理
}
def ask(task: str, messages: list, **kw):
return completion(model=MODEL_MAP[task], messages=messages, **kw)
把「换模型」抽象成配置之外,还要准备降级链:主模型超时或报错时自动切备用;并用「影子流量」在小比例真实请求上并行跑候选模型,离线比对质量与成本,再灰度切换。
python# 带降级与成本记账的路由(生产骨架,依赖 litellm)
import time
from litellm import completion
CHAIN = { # 每个任务一条降级链
"reasoning": ["claude-opus-4.7", "gpt-5", "deepseek-v4"],
"bulk": ["deepseek-v4", "qwen3.5-235b"],
}
SPEND = {} # 累计花费(USD)
def ask(task, messages, **kw):
for model in CHAIN[task]:
try:
t0 = time.time()
r = completion(model=model, messages=messages, timeout=60, **kw)
SPEND[model] = SPEND.get(model, 0) + r.usage.total_tokens * 1e-6
return r, model, round(time.time() - t0, 2)
except Exception as e:
print(f"{model} 失败,降级:{e}")
raise RuntimeError("全部模型失败")
# 预期:主模型异常时自动切下一档;SPEND 可按模型核对预算
# 关键:把「模型选择」收敛为可观测、可回滚的一处配置
- 做决策:给「每天 500 万次短分类」选一个模型并写出理由。判据:以成本与一致性为先,选低价 MoE(如 DeepSeek-V4 / Qwen3.5),而非旗舰。
- 算总成本:某 API 每百万 token 0.28 USD,日耗 4 亿 token,月成本?参考答案:约 0.28×400×30 ≈ 3360 USD。
- 风险点:列出三条「只看排行榜选型」的失败模式。判据:能力维度错配、成本失控、供应商锁定。
- 工程抽象:如何做到换模型不改业务代码?判据:统一网关 / provider 抽象层(LiteLLM),模型名进配置。
- 进阶:设计一个「影子流量」灰度方案。判据:小比例真实请求并行跑候选模型,离线比对质量与成本后再切主流量。
4. 从零实现 MiniGPT
学习路径
- 读 4.1:逐行理解 MiniGPT 结构(Token/Pos Embedding、RMSNorm+Pre-LN、GQA、交叉熵、Top-K/温度采样)
- 亲手实现或敲一遍 MiniGPT,回答"删除某行会怎样"并测生成
- 对接 M8:以该实现为基线逐层扩展出可加载开源权重的 GPT
核心知识点详解
- MiniGPT 的骨架顺序:Token Embedding + Pos Embedding → RMSNorm+Pre-LN → GQA 注意力 → SwiGLU FFN → 最后一个 LN → LM head。删掉任何一块 Pre-LN / 残差 / 归一化都要能说出它的作用。
- 交叉熵是预训练目标:对每个位置,用 softmax over vocab 预测下一个 token,
F.cross_entropy(logits, labels)只对目标 id 回传。nn.CrossEntropyLoss自带 softmax 与 NLL 合并,别自己再加一层 log_softmax。 - Top-K / 温度采样生成:温度 τ=0.7 缩放 logits 后做
top_k截断再 softmax 采样;τ→0 退化为贪心,τ 越大越多样。生成要记得裁剪上下文长度并按 batch 并行,否则会越跑越慢。 - 验证每一行的办法:把某行注释掉再跑生成,对比 PPL 与输出质量即可验证该行必要性。常见坑:把 embedding dim 与 vocab 形状写反,
nn.Embedding(vocab, d)第一个参数是词表大小。
学习路径
- 读 4.1:背下 embedding + 每块 ≈ 12d²、权重共享、std=0.02、输出投影零初始化
- 用代码核对参数量估算,验证 weight tying 与残差缩放 1/√N
- 对接 M8:按此初始化方案写模型,确保与开源权重加载后行为一致
核心知识点详解
- 总参数的快速估算:总参 ≈ embedding(vocab·d)+ N·每块(12d²),含 weight tying 时 output head 复用 token embedding,省去 d·vocab。用
sum(p.numel())核对,估算应和实际一致到 1–2% 量级。 - std=0.02 初始化:层权重用 N(0, 0.02²) 初始化,控制前向方差随深度增长的机制;残差缩放 1/√N(N 为层数)进一步抑制深层激活爆炸。
- 输出投影零初始化:FFN 的输出投影(down_proj)用 0 初始,使残差流在训练起点不被打乱,模型从「接近 identity」的稳定态起步,深度学习常用技巧。常见坑:weight tying 与单独 output head 二选一,别同时定义导致参数量与开源权重对不上。
- 对齐开源权重的初始化陷阱:若加载开源权重后 logits 差一个量级,先查初始化分布、
eps、bias开关与 RMSNorm 实现是否一致,再谈训练稳定。
学习路径
- 读 4.1:理解对齐开源权重需复现 tokenizer、RoPE base、eps 与权重共享
- 加载开源权重并用 logits 对比,让逐层输出与 HuggingFace 对齐
- 对接 M8:完成 test_parity,使 max abs diff < 1e-3(cos > 0.9999)
核心知识点详解
- 对齐的四要素:tokenizer / base / eps / 权重共享:要输出与 HuggingFace 一致,需复现字节级 tokenizer(带 BOS/EOS)、RoPE base、RMSNorm 的 eps、以及是否 weight tying。任一不一致都会让 logits 累积偏差。
- 逐层对齐而不是只看最后一层:用
state_dict逐层加载后,先对中间层输出做 cos 对比定位第一个偏差层,再逐步修。final logits 的max abs diff < 1e-3、cos > 0.9999才是通过判据。 - 数值类型与 kernel 也会造成偏差:bf16 vs fp32 累积、FlashAttention vs math kernel 都可能带来 1e-3 级差异。固定用同一 dtype 与 attention 实现再比对才公平。常见坑:把 drop 层默认 dropout>0 没关,推理时输出随机抖动。
4.1 完整可训练实现
下面这段代码是一个可以真正训练出「会说话」的字符级 GPT 的最小完整实现。建议亲手敲一遍,然后逐行回答「这一行删掉会怎样」。
pythonimport torch, torch.nn as nn, torch.nn.functional as F
class MiniGPT(nn.Module):
def __init__(self, vocab=65, d=256, layers=6, heads=8, max_len=256):
super().__init__()
self.tok_emb = nn.Embedding(vocab, d)
self.pos_emb = nn.Embedding(max_len, d)
self.blocks = nn.ModuleList([Block(d, heads, g=heads // 2) for _ in range(layers)])
self.norm = RMSNorm(d)
self.head = nn.Linear(d, vocab, bias=False)
self.apply(self._init)
@staticmethod
def _init(m):
if isinstance(m, nn.Linear):
nn.init.normal_(m.weight, std=0.02)
elif isinstance(m, nn.Embedding):
nn.init.normal_(m.weight, std=0.02)
def forward(self, idx, targets=None):
B, T = idx.shape
x = self.tok_emb(idx) + self.pos_emb(torch.arange(T, device=idx.device))
for blk in self.blocks:
x = blk(x)
logits = self.head(self.norm(x))
if targets is None:
return logits, None
loss = F.cross_entropy(logits.view(-1, logits.size(-1)),
targets.reshape(-1), ignore_index=-1)
return logits, loss
@torch.no_grad()
def generate(self, idx, max_new=200, temperature=0.8, top_k=50):
for _ in range(max_new):
idx_cond = idx[:, -256:] # 上下文裁剪
logits, _ = self(idx_cond)
logits = logits[:, -1, :] / temperature
if top_k: # Top-K 采样
v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
logits[logits < v[:, [-1]]] = -float("inf")
probs = F.softmax(logits, dim=-1)
idx = torch.cat([idx, torch.multinomial(probs, 1)], dim=1)
return idx
# 训练骨架(真实训练要加 AMP、梯度累积、余弦调度、验证集)
model = MiniGPT().cuda()
opt = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=0.1)
for step, (x, y) in enumerate(loader):
_, loss = model(x, y)
opt.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
if step % 100 == 0:
print(f"step {step} loss {loss.item():.4f} ppl {torch.exp(loss).item():.2f}")
从零实现一个能对齐开源权重的 GPT,第一步是把参数量算清楚。一个标准 Decoder-only GPT 的参数由四块构成:① token embedding(词表 V × d_model);② 每个块的注意力(4 个投影 4·d²,若 QKV 合并为一个 3d² 的矩阵再加 Out d²,仍是 4d²);③ 每个块的 FFN(SwiGLU 时为 3·d·d_ff,d_ff=8/3·d 时约 8d²);④ 最终的 RMSNorm(2d,权重)与 LM head。关键细节:当 embedding 与 LM head 权重共享(weight tying)时,二者不重复计入;若不用位置嵌入(用 RoPE)则没有 pos_embedding 参数。
python# 估算一个 GPT 的参数量(含权重共享)
def gpt_params(V, d, L, tied=True, ffn_mult=8/3):
emb = V * d # token embedding
per_block = 4*d*d + 3*d*int(d*ffn_mult) # attn(4d^2) + SwiGLU FFN(3*d*8/3 d)
blocks = L * per_block
final_norm = 2 * d # RMSNorm 权重(无 bias)
lm_head = 0 if tied else V * d # 权重共享则不重复计
return emb + blocks + final_norm + lm_head
# 例:d=4096, L=32, V=100000, tied=True
P = gpt_params(100000, 4096, 32, tied=True)
print(f"总参数 ≈ {P/1e9:.2f}B") # 约 1.24B
print("embedding:", round(100000*4096/1e9,3), "B 每 block:", round((4*4096*4096 + 3*4096*int(4096*8/3))/1e6,1), "M")
初始化尺度同样决定成败:embedding 与线性层权重常用 std = 0.02 的正态初始化(与 GPT-2 一致);注意力输出投影常零初始化,使每个新块的初始输出≈0,从而初始时整个网络近似恒等映射、残差主路不被破坏(这也是深模型能稳定起训的技巧之一)。2026 前沿在此基础上叠加:残差缩放(每个残差分支乘 1/√N 或 ReZero 标量)、QK-Norm(对 Q、K 再各做一次归一化)、以及把优化器换成 Muon(对 2D 权重做谱范数正交化,配合 AdamW 用于 1D 向量)以获得更稳更大的更新。
python# 对齐开源权重后:用余弦相似度与最大绝对误差逐层校验
import torch, torch.nn.functional as F
@torch.no_grad()
def compare(model, ref_model, tokens, atol=1e-3):
out_m = model(tokens)[0].float()
out_r = ref_model(tokens)[0].float()
cos = F.cosine_similarity(out_m.flatten(), out_r.flatten(), dim=0).item()
err = (out_m - out_r).abs().max().item()
print(f"cos={cos:.6f} max_abs_err={err:.2e}")
return cos > 0.9999 and err < atol
# 预期:对齐正确时 cos≈1.000000、max_abs_err < 1e-3(BF16 下放宽到 3e-3)
# 若 cos 明显小于 1:优先排查 tokenizer、RoPE base、RMSNorm eps、权重共享
把参数量算清楚是「从零复现」的第一道关。以 d=4096、L=32、V=100k、权重共享为例:embedding 约 0.41B、每块约 12d²≈201M、32 块≈6.4B,合计落在 7B 量级;把 L 提到 80 就接近 70B 级。手算一遍,你才会对「d_model、层数、词表」三个旋钮如何共同决定模型大小形成直觉。
- 跑通:让 MiniGPT 在莎士比亚字符语料上训到 PPL < 3(约几千步、单卡即可)。判据:采样输出有英文单词与标点结构。
- 删行实验:删掉 Pre-LN 的 RMSNorm 或残差,观察 loss 是否发散。判据:删残差必定发散;删 RMSNorm 在深层不稳。
- 算参数量:V=100k、d=4096、L=32、权重共享,估算总参。参考答案:约 7B 量级(含 embedding 0.41B)。
- 对齐:用固定输入对比自实现与官方权重的 logits。判据:cos > 0.9999、max_abs_err < 1e-3。
- 采样:对比 temperature=0.2 与 1.0、top_k=1 与 50 的输出差异。判据:低温/小 top_k 更确定易重复;高温更发散。
5. Scaling Law 与规模化的现实
学习路径
- 读 5.1:从 Kaplan→Chinchilla(20 tokens/param)→2026 过训练(100+)理解分配演进
- 用 total_cost 函数比较 70B/13B/7B 的训练与长期推理成本,理解过训练取舍
- 对接 M8:在 architecture.md 用参数量/FLOPs 拆解与 C≈6ND 估算串起对应关系
核心知识点详解
- Chinchilla 的 20 tokens/param 最优分配:约 20 tokens/param 附近让「计算量固定下的 PPL 最低」;例如 7B 模型配 140B tokens(7B×20)。这是计算最优,不是数据最优——数据通常更充裕。
- 2026 走向「过训练」100+ tokens/param:Qwen3 约 36T tokens,远超 20/param 的 Chinchilla 线。动机:推理多、训练少时,把小模型多过训练在长期总量成本上更划算,牺牲一点训练 FLOPs 效率换取推理更便宜。
- 用 total_cost 比较长期成本:把训练成本 + 持续推理成本都算进去,70B/13B/7B 三个规模在「训多少、跑多久」坐标下趋势会反转。常见坑:只看训练成本忽略长期推理 token 量,导致规模选错。
- 数据质量 > 数量:同样 token 下,去重后高质量数据的 PPL 下降明显优于低质重复数据。清出来的预算值得优先花在高质量来源上。
学习路径
- 读 5.2:背下 C≈6ND、MFU 30–50%、每参数约 18 字节显存、ZeRO-3 每卡约 5 字节
- 用 C≈6ND 估算 70B/15T 算力量级并粗算 2048 卡的训练天数
- 对接 M8:在 architecture.md 给出模型 FLOPs 与显存估算,供手写实现核对规模
核心知识点详解
- C ≈ 6ND 是算力基准式:N 参数、D tokens 时,一次前向+反向的 FLOPs ≈ 6ND(另有注意力等常量阶修正)。70B/15T:6×7e10×1.5e13 ≈ 6.3e24 FLOPs。
- MFU 30–50% 决定墙钟时间:MFU=实际/峰值利用率。训练天数 = C /(卡数×单卡峰值 FLOPs×MFU×3600×24)。按此粗算 2048 卡 H100 训 70B/15T 约 70–90 天。常见坑:分母忘了乘 MFU,把理论峰值当实际收益,天数严重低估。
- 显存分摊与并行策略:权重+梯度+Adam 状态等每参数约 18 字节(fp16/bf16 + 主权重 fp32 + 一阶/二阶动量);ZeRO-3 把冗余状态切走,每卡约 5 字节。DP 复制数据、TP 切张量维、PP 切层、EP 切专家、SP 切序列轴。
- loss spike 回滚:训练崩溃前通常有 loss 尖峰预警,回滚到上一个稳定 checkpoint、降 lr(如 −50%)再继续,而不是硬扛。
核心知识点详解
- 验证 70B/15T ≈ 6.3e24:6×7e10×1.5e13=6.3e24 FLOPs。这是「规模→算力」的一步换算,面试 / 规划都常用。
- 2048 卡约 90 天:天数 = 6.3e24/(2048×峰值×MFU)。套 MFU≈40% 即得 90 天量级。
- 7B 按 20 tokens/param 配 140B:7e9×20=1.4e11 = 140B tokens,即 Chinchilla 最优分配的最简洁应用。
5.1 从 Kaplan 到 Chinchilla 再到 2026 的实践
Scaling Law 回答的问题是:给定算力预算,参数与数据应该怎么分配。Kaplan(2020)认为应该优先放大模型;Chinchilla(2022)修正为「参数与数据应大致同比例增长,约 20 token/参数」;而 2026 年的实际做法是刻意过训练——用远多于 Chinchilla 比例的数据训一个较小的模型,因为推理成本才是长期成本,小模型服务更便宜。
| 阶段 | 核心结论 | 生产含义 |
|---|---|---|
| Kaplan 2020 | loss 随参数量、数据量、算力呈幂律下降 | 开启了“越大越好”的军备竞赛 |
| Chinchilla 2022 | 给定算力,参数与数据应同比例增长(≈20 tokens/param) | 纠正了“只堆参数不看数据”的浪费 |
| 2026 实践 | 刻意过训练小模型(100+ tokens/param),能力接近大模型但推理便宜 | 端侧与高并发场景的主流选择 |
| 数据质量优先 | 去重、筛选、配比、合成数据的收益常大于规模 | 数据工程成为核心竞争力 |
| 推理时算力 | 预训练收益递减,把算力投到推理期(长思考、多次采样) | 催生推理模型与 test-time scaling |
Chinchilla 与「过训练」的差别可用一句话的量级表达:给定推理请求量,过训练把更多算力前置到训练、换取更低的长期推理成本——当推理总 token 数远超训练 token 数时,小模型过训练的总成本更优。
python# 「过训练」何时划算:训练成本 vs 长期推理成本
def total_cost(model_b, train_tokens, serve_tokens, c_train=6e-12, c_serve=2e-12):
train = c_train * 6 * model_b * 1e9 * train_tokens # 6ND
serve = c_serve * 2 * model_b * 1e9 * serve_tokens # 2ND_serve
return train, serve
for model_b, tr_tok in [(70, 1.4e12), (13, 5e12), (7, 1e13)]:
tr, sv = total_cost(model_b, tr_tok, serve_tokens=5e15)
print(f"{model_b:3d}B 训练={tr/1e18:.2f} 推理={sv/1e18:.2f} EFLOPs 合计={(tr+sv)/1e18:.2f}")
# 预期:大模型训练与推理都贵;小模型过训练后推理成本低,长期更省
# 这就是 2026 端侧/高并发普遍选「小模型 + 100+ tokens/param」的算力理由
| 代表(2026) | 训练 token 量级 | 备注 |
|---|---|---|
| Qwen3 系列 | 约 36T tokens | 刻意过训练,小杯也吃大量 token |
| DeepSeek-V3/V4 | 约 14.8T 及以上 | MoE + FP8,token/param 高于 Chinchilla |
| Llama 4 | 数十 T | 多模态语料并入 + 长上下文续训 |
| 开源小模型(1B–8B) | 数 T 到数十 T | 过训练以适配端侧与高并发 |
5.2 训练并行与算力估算
| 并行方式 | 切分对象 | 解决的问题 | 代价 |
|---|---|---|---|
| 数据并行 DP | 样本 | 吞吐 | 需同步梯度,通信量大 |
| 张量并行 TP | 单层内的矩阵 | 单层放不下 | 每层都要通信,需高带宽互联(NVLink) |
| 流水线并行 PP | 按层切分 | 模型太深 | 有气泡(bubble),需要微批次调度 |
| 序列并行 SP | 序列维度 | 长序列激活过大 | 与 TP 配合使用 |
| 专家并行 EP | MoE 专家 | MoE 参数过大 | token dispatch/combine 通信 |
| ZeRO / FSDP | 优化器状态、梯度、参数 | 显存 | 通信换显存,重叠做得好才划算 |
python# 训练算力估算:C ≈ 6 · N · D (N=参数量, D=训练 token 数)
# 6 = 前向 2 + 反向 4;这是规划预算与集群规模的第一性公式
def training_flops(n_params, n_tokens):
return 6 * n_params * n_tokens
def gpu_hours(flops, gpu_tflops=990, mfu=0.4):
"""mfu: 模型算力利用率,大集群通常 30%-50%"""
return flops / (gpu_tflops * 1e12 * mfu) / 3600
N, D = 70e9, 15e12 # 70B 参数,15T tokens
f = training_flops(N, D)
print(f"总算力 ≈ {f:.2e} FLOPs")
print(f"H100 集群所需 ≈ {gpu_hours(f):,.0f} GPU·小时")
print(f"2048 卡集群 ≈ {gpu_hours(f)/2048/24:.1f} 天")
MFU(模型算力利用率)是集群健康度的第一指标:大集群 30%–50% 属正常,低于 20% 说明通信或数据加载在拖后腿。ZeRO 阶段决定显存与通信的权衡:ZeRO-1 只切优化器状态、ZeRO-2 再切梯度、ZeRO-3(≈ FSDP)连参数都切,显存降到约 1/N,但通信量逐级上升。
python# 混合精度训练:每参数显存开销与 ZeRO 阶段的收益
def mem_per_param(dtype_bytes=2, zero=0, n_gpu=8):
base = dtype_bytes + dtype_bytes # 参数 + 梯度(bf16)
opt = 12 if zero < 1 else 12 / n_gpu # Adam 约 12 字节/参数
if zero >= 2: base = base / n_gpu # 梯度分片
share = dtype_bytes if zero < 3 else dtype_bytes / n_gpu
return base + opt + share
for z in (0, 1, 2, 3):
print(f"ZeRO-{z} (8 卡): 每参数约 {mem_per_param(zero=z):5.2f} 字节")
# 预期:ZeRO-0 约 18 字节/参数,ZeRO-3 降到约 5 字节/参数
# 量级:70B 模型 ZeRO-3 常驻约 350GB,需 8x80GB 卡才放得下参数 + 优化器
- 算算力:70B 参数、15T token,C ≈ 6ND 是多少?参考答案:6×70e9×15e12 ≈ 6.3e24 FLOPs。
- 算时间:2048 张 H100(峰值 989 TFLOPS、MFU 40%)需多久?参考答案:6.3e24/(2048×989e12×0.4) ≈ 7.8e6 秒 ≈ 90 天。
- Chinchilla 判断:7B 模型按 20 tokens/param 应配多少数据?参考答案:约 140B token。
- 过训练取舍:为什么 2026 常用 100+ tokens/param?判据:推理成本主导长期成本,小模型服务更便宜。
- 显存估算:混合精度训练每参数显存约多少?判据:参数+梯度+Adam 状态约 18 字节/参数(未分片)。
6. FFN、归一化与残差:训练稳定性的根基
学习路径
- 读 6.1:理解 FFN 占总参数 2/3,SwiGLU 用 hidden=8/3·d 保持与 4d MLP 等价参数量
- 用代码跑参数量拆解,验算 FFN:注意力 = 8d²:4d² = 2:1
- 对接 M8:手写 FFN 时采用 SwiGLU,保持总参数与参考模型一致
核心知识点详解
- FFN 占总参数 2/3:ReLU FFN:d×4d(上投影)+4d×d(下投影)= 8d²;注意力投影 4d²。故 FFN:注意力 = 8d²:4d² = 2:1,FFN 占全部(含自身)约 2/3。
- SwiGLU 用 hidden=8/3·d 保持等价参数量:SwiGLU 是两个分支(gate+up)+ 输出三组投影,若仍用 4d 会膨胀。取中间宽=8/3·d:3×8/3d=8d,总投影 ≈ 8/3d×d×3 + 8/3d×d ≈ 8d²,与 4d MLP 等价。公式:
up=gate=down=8*d//3。 - 为什么门控更好:
SwiGLU(x)=x·W1 · σ(x·WGo)⊗(x·Wu),用可学习的门控选择信息,比 ReLU 的非线性更有表达力,训练更稳。门控函数是F.silu(=swish)。 - MoE 是 FFN 的稀疏改造:把单个大 FFN 换成 N 个专家 FFN + router,token 只走 Top-K 个,激活参 < 总参。常见坑:hidden 取整要能整除分组(如 d 是 8/3 倍数),否则分组 padding 或维度对不上权重。
学习路径
- 读 6.2:理解 Post-LN→Pre-LN 的梯度直通、RMSNorm 只除均方根、残差缩放 ReZero、QK-Norm
- 对比 Pre-LN vs Post-LN 是否需 warmup,验证 1–3% 步数的经验值
- 对接 M8:手写 GPT 采用 Pre-LN + RMSNorm + 残差,配置 warmup
核心知识点详解
- Post-LN → Pre-LN 的动机:Post-LN(归一化在残差之后)深层梯度易爆;Pre-LN(归一化在子层之前)让残差流恒等直通,梯度 1 阶传播,深层更稳。2026 主流几乎全 Pre-LN。
- RMSNorm 只除均方根:
x / sqrt(mean(x²)+eps)·γ,无均值中心化与平移偏置,比 LayerNorm 少几次归约、且可被后续线性层吸收。常见坑:归一化的eps归属不同(RMS 里 vs 分母里)会造成 1e-3 级输出偏差。 - warmup 占总步数 1–3%:Pre-LN 相对 Post-LN 对 warmup 不那么敏感,但经验上仍配 1–3% 步数的线性 warmup 稳训练。
min(t/T_warm)·lr_max即可实现。 - ReZero 残差缩放与 QK-Norm:ReZero 用可学习缩放 α 初始≈0 控制残差贡献;QK-Norm 对 Q/K 先归一化防熵坍缩。两者都是「初始化稳定」路线。
学习路径
- 读 6.3:理解 Muon 谱范数正交化、用于 2D 权重而 AdamW 处理 1D,DeepSeek-V4 迁移
- 在一个小模型上对比 AdamW 与 Muon 的收敛与吞吐,确认 Muon 的边界
- 完成本章自测:说明 Muon 的 weight decay 为何不能直接用 AdamW 数值
核心知识点详解
- Muon 用谱范数正交化更新 2D 权重:对参数矩阵求梯度后用「牛顿-舒尔兹迭代」更新,使权重的奇异值谱趋向均衡、抑制条件数爆炸,从而放宽对学习率的敏感性。
- AdamW 处理 1D 参数、Muon 处理 2D:embedding/output 等 1D 参数仍走 AdamW;卷积/线性等 2D 权重走 Muon(按行或按列分别正交化)。混合使用是 DeepSeek-V4 迁移的标准配方。
- weight decay 不能直接沿用 AdamW 数值:Muon 的更新几何不同,若照抄 AdamW 的
weight_decay=0.1会过正则。需专门校调 wd。常见坑:把 AdamW 整组超参一把搬到 Muon,收敛反而变差。 - 适用边界:Muon 主要利好深、宽的全连接/卷积主干;小模型收益不确定,需在小规模上 A/B 对比收敛与吞吐再上量。
6.1 前馈层与 SwiGLU:参数量的真正大头
一个 Transformer 块的参数分布极不均匀:FFN 通常占全部参数的约 2/3。原因是注意力只有 4 个投影矩阵(Q/K/V/Out,各 d_model×d_model),而 FFN 是两个更大的矩阵。以标准 FFN(隐藏维度 d_ff = 4·d_model、无 bias)为例:W₁ 参数量 d_model·d_ff = 4d²,W₂ 参数量 d_ff·d_model = 4d²,合计 8d²;而注意力 QKV+Out = 4d²。所以 FFN : 注意力 = 8 : 4 = 2 : 1,即 FFN 占 2/3。这也解释了为什么 MoE 选择改造 FFN——那里藏着最多的「闲置容量」可供稀疏激活。
python# 参数量拆解:为什么 FFN 是 2/3
d = 4096 # d_model
d_ff = 4 * d # 经典隐藏维度
attn_params = 4 * d * d # Q,K,V,Out 四个投影
ffn_params = 2 * d * d_ff # W1, W2
total = attn_params + ffn_params
print(f"attention = {attn_params:,} ffn = {ffn_params:,}")
print(f"FFN 占比 = {ffn_params/total:.1%}") # 0.667
原始 Transformer 用 ReLU FFN:FFN(x) = W₂·max(0, W₁x)。2020 年后 GELU / GeGLU / SwiGLU 成为主流。SwiGLU 用一个额外的「门」把两个线性变换按元素相乘:FFN(x) = (W₁x) ⊙ SiLU(W_g·x),再经 W₂ 投影。因为需要两个「向上投影」矩阵(W₁ 与 W_g),为了保持与 4d FFN 相近的参数量,实际隐藏维度取 8/3·d_model ≈ 2.67·d_model 而非 4d。GELU / SiLU 相比 ReLU 在原点附近更平滑、非零梯度区间更大,配合 RMSNorm 在大模型上更稳定。
| FFN 变体 | 公式(核心) | 参数量(无 bias) | 备注 |
|---|---|---|---|
| ReLU FFN | W₂·ReLU(W₁x) | 2·d·d_ff | 原版,d_ff=4d 最常见 |
| GELU FFN | W₂·GELU(W₁x) | 2·d·d_ff | BERT / 早期 GPT |
| GeGLU | W₂·(W₁x ⊙ GELU(W_g x)) | 3·d·d_ff | 用 GELU 做门 |
| SwiGLU | W₂·(W₁x ⊙ SiLU(W_g x)) | 3·d·d_ff (d_ff=8/3 d) | LLaMA / Qwen / DeepSeek 标配 |
| SwiGLU (d_ff=4d) | 同上 | 3·d·4d = 12d² | 容量更大但更贵 |
python# SwiGLU 的矩阵形状(与标准 FFN 对比)
import torch, torch.nn as nn
d, mult = 4096, 8/3
hid = int(d * mult) # ≈ 10922,比 4d 小
sg = nn.Sequential(
nn.Linear(d, hid, bias=False), # W1
nn.SiLU(),
nn.Linear(d, hid, bias=False), # Wg(门)
)
out = nn.Linear(hid, d, bias=False) # W2
# 注意:SwiGLU 需要把 W1 与 Wg 的输出逐元素相乘,故 hid 要小于 4d
# 才能保持与 ReLU-FFN(4d) 相当的参数量:3*hid*d = 2*(4d)*d -> hid = 8/3 d
python# 三种 FFN 的参数量对比(d=4096,单个块,无 bias)
d = 4096
relu = 2 * d * (4 * d) # W1, W2;d_ff = 4d
swiglu = 3 * d * int(d * 8 / 3) # W1, Wg, W2;d_ff = 8/3 d
print(f"ReLU FFN = {relu/1e6:7.1f} M")
print(f"SwiGLU FFN = {swiglu/1e6:7.1f} M (与 ReLU 基本持平)")
# 预期:两者都在 134M 量级 —— SwiGLU 用 8/3 d 的隐藏维换来门控表达力,却不增加参数
6.2 归一化与残差:为什么深模型需要 warmup
归一化的位置是 2017 原版与现代块最大的结构差异之一。Post-LN(原版)把 LayerNorm 放在残差相加之后:x_{out} = LN(x + Sublayer(x))。Pre-LN 把 LN 放在子层输入之前:x_{out} = x + Sublayer(LN(x))。差别看似微小,却决定了深层能否训稳:Post-LN 的梯度要穿过层层 LN 才能回到浅层,深层时梯度方差随深度剧烈变化,必须配合极精细的 warmup;Pre-LN 在残差主路上没有任何归一化,形成一条「干净」的恒等通路,梯度可直接回传,对 warmup 不敏感——这正是现代大模型默认 Pre-LN 的根因。
| 维度 | Post-LN | Pre-LN |
|---|---|---|
| 梯度传播 | 需穿过每层 LN,深层易不稳 | 残差主路无 LN,梯度直通 |
| warmup 需求 | 必须,且要长而精细 | 弱得多,可省或很短 |
| 训练动态 | 初期易 loss spike | 更稳定 |
| 最终性能 | 充分训练后略优 | 与现代配方等价,工程更稳 |
| 代表 | 原版 Transformer | LLaMA / Qwen / GPT 现代块 |
RMSNorm 是 LayerNorm 的减负版:只除以激活的均方根、不做均值中心化、去掉偏置项。公式:x̂ = x / √(mean(x²) + ε) · γ。它比 LayerNorm 少一次均值减法与一次仿射,在长序列上省下的算力可观,且效果相当,因此被 LLaMA / Qwen / DeepSeek 等几乎全部采用。
残差连接带来一个微妙问题:每一层都把子层输出加回主路,若子层输出方差不控,深层激活的方差会逐层累积、爆炸或坍缩。现代实践用三种手段压制:① 残差缩放(如初始把每个残差分支乘 1/√N,或 ReZero 把分支乘可学习标量初始化为 0);② QK-Norm(在注意力内部对 Q、K 再各做一次归一化,DeepSeek / 稳定训练常用);③ 初始化尺度(Embedding 与残差初始 std ≈ 0.02,配合 warmup 让早期注意力接近均匀、避免突然聚焦到噪声)。这三点正是「为什么深模型一上来就用大学习率就会炸、必须先 warmup」的定量答案。
python# 残差缩放在 2026 多模态 / 深模型里很常见(概念示意)
class BlockWithScale(nn.Module):
def __init__(self, d, h, g):
super().__init__()
self.n1, self.attn = RMSNorm(d), GQAttention(d, h, g)
self.n2, self.ff = RMSNorm(d), SwiGLUFFN(d)
self.alpha = nn.Parameter(torch.zeros(1)) # ReZero:初始为 0,分支先「关闭」
def forward(self, x):
x = x + self.alpha * self.attn(self.n1(x))
x = x + self.alpha * self.ff(self.n2(x))
return x
# 训练早期 alpha≈0 => 网络≈单纯恒等映射,梯度干净;
# 随训练进行 alpha 被解锁,深度逐步「长出来」,这是对 warmup 的另一种解法。
python# 学习率 warmup:为什么深模型不能一上来就用大 lr
import math
def warmup_cosine(step, total, warmup_frac=0.015, base_lr=3e-4, min_lr=3e-5):
w = int(total * warmup_frac) # warmup 约占总步数 1%-3%
if step < w:
return base_lr * (step + 1) / w # 线性升
p = (step - w) / max(1, total - w) # 余弦衰减
return min_lr + 0.5 * (base_lr - min_lr) * (1 + math.cos(math.pi * p))
for s in (0, 100, 300, 5000, 20000):
print(f"step {s:6d} lr={warmup_cosine(s, 20000):.2e}")
# 预期:前约 300 步线性上升到 3e-4,随后余弦衰减到 3e-5
# 经验:warmup 占总步数 1%-3%;过短易 loss spike,过长浪费算力
2026 稳定化再补一条:Muon 优化器对 2D 权重做 Newton-Schulz 迭代的谱范数正交化,使更新矩阵奇异值更均匀,并配合 AdamW 处理 1D 参数(RMSNorm 权重、bias)。DeepSeek-V4 报告 Muon 在同 token 预算下 loss 更低、对大学习率更鲁棒,是 2026 从 AdamW 迁移的主要方向之一。
- 算比例:d=4096、ReLU FFN(d_ff=4d)、注意力 4d²,FFN 占比?参考答案:8d²/(8d²+4d²)=2/3。
- 解释 SwiGLU:为什么隐藏维取 8/3 d 而不是 4d?判据:门控需两个上投影(3 个矩阵),8/3 d 使参数量与 2·4d 持平。
- warmup 取舍:warmup 太长/太短的后果?判据:太短初期 loss spike,太长浪费算力;常取总步数 1%-3%。
- Pre-LN 论证:为什么 Pre-LN 对 warmup 不敏感?判据:残差主路无归一化,梯度可直通浅层。
- Muon 判断:它和 AdamW 分别用在哪些参数上?判据:Muon 用于 2D 权重做正交化,AdamW 用于 1D 参数。
7. KV Cache 与推理成本工程
学习路径
- 读 7.1:默写 KV 显存公式 KV=2·L·kv_heads·head_dim·seq·batch·dtype,理解每步降到 O(T)
- 用 kv_cache_gb 函数实测常见配置下的显存,理解 memory-bound 与量化/分页
- 对接 M8:实现 kv_cache.py 并写出增量解码,与 M2 显存公式互验
核心知识点详解
- KV 显存公式:KV_bytes = 2 × L × kv_heads × head_dim × seq × batch × dtype_bytes。系数 2 是 K 与 V 各一份。L=层数、head_dim=每头维度。
- 预填充 O(T²) → 解码 O(T):预填充并行算整个序列(KV 全存),解码每步只算 1 个 token 但要把 KV Cache 整体搬进带宽,复杂度降到 O(T),显存仍随 T 线性增长。
- 省显存的三种手段:量子化 KV(FP8 减半、INT4 到 1/4)、PagedAttention 按页分配消除碎片与预留浪费、GQA/MQA 砍 kv_heads。三者可叠加,量化需注意精度回退。
- memory-bound 的含义:解码阶段算术强度≈低(每个 token 的 FLOPs 少而显存读写多),瓶颈在带宽而非计算,因此省 KV 直接换吞吐。常见坑:128K 上下文仍按小模型 bf16 估算,实际 KV 已占掉大半显存。
学习路径
- 读 7.2:掌握每 token 约 0.3125 MB、32K 约 10.2 GB、128K 约 40.96 GB 的量级
- 手算 128K 上下文 40.96 GB 并核验 MHA 是 GQA 8 倍
- 对接 M8:用 KV 显存公式做并发容量规划,指导 kv_cache 实现的 batch/seq 上限
核心知识点详解
- 每 token 约 0.3125 MB:取 L=40、kv_heads=8、head_dim=128、bf16(2B):2×40×8×128×2 = 327,680 B ≈ 0.3125 MB/token。单位记忆值够用于秒估。
- 32K 约 10.2 GB、128K 约 40.96 GB:0.3125MB×32768≈10.2GB;×131072≈40.96GB。128K 时单条请求就把常规 40GB 卡吃掉大半,必须量化或换大卡。
- MHA 是 GQA 的 8 倍:MHA 的 kv_heads 与 query 头数相同(h=64),GQA 用 8 个共享,故 KV 显存比为 64:8=8:1。验证:二者只在这个系数上不同。
- 并发容量规划:单 GPU 可用显存 / (每并发KV+权重+激活) = 并发上限。常见坑:忘了预留权重与激活显存,只按 KV 规划并发导致推理阶段 OOM。
kv_cache_gb结合get_available_gb直接验证。
核心知识点详解
- 默写公式与速算:KV=2·L·kv_heads·head_dim·seq·batch·dtype;配合「每 token≈0.3125MB、32K≈10.2GB、128K≈40.96GB」两档速记即可独立复现。
- 128K ≈ 40.96 GB 的推导:0.3125MB×131072 = 40.96GB,且 MHA 为该值的 8 倍即 ≈327GB(接近超长上下文要 CP 拆的原因)。
- PagedAttention 为何提吞吐:以固定页块分配连续显存、按需分配消除内部碎片与请求间预留浪费,提升批次可装下的并发,从而提吞吐。常见坑:分页仍有页内平均 ~50% 碎片,量化与 GQA 才是硬减。
7.1 KV Cache 原理与显存公式
自回归生成时,第 t 步的注意力需要第 1…t 步所有 token 的 K、V。若每步都重新计算,复杂度是 O(T²)。KV Cache 把这些历史 K、V 缓存起来,使每生成一个新 token 只需算当前 query 与已存 K/V 的点积,单步降到 O(T)。代价是显存要随上下文线性增长——这正是长上下文推理的「内存墙」。
显存公式(每 batch、每序列):KV_bytes = 2 × layers × kv_heads × head_dim × seq_len × batch × dtype_bytes。其中 2 来自 K 与 V 各一份;kv_heads 是 KV 头数(GQA 下远小于注意力头数);dtype_bytes 取决于精度(FP16/BF16 = 2,FP8/INT8 = 1,FP4/INT4 = 0.5)。注意它只与序列长度和层数线性增长,与 batch 成线性关系,因此大并发长上下文场景的显存主要由 KV Cache 主导。
| dtype | 字节/元素 | 备注 |
|---|---|---|
| FP32 | 4 | 几乎不用作 KV |
| FP16 | 2 | 常见基线 |
| BF16 | 2 | 训练友好,推理常用 |
| FP8 (E4M3/E5M2) | 1 | 2026 大模型推理常态 |
| INT8 | 1 | 量化推理 |
| FP4 / INT4 | 0.5 | 激进压缩,需校准 |
python# KV Cache 显存(自包含版本,便于直接复用)
def kv_cache_gb(layers, seq_len, kv_heads, head_dim, batch=1, dtype_bytes=2):
bytes_ = 2 * layers * kv_heads * head_dim * seq_len * batch * dtype_bytes
return bytes_ / 1024**3
# 形状备忘:每个 token 每层存 K 与 V 各 (kv_heads, head_dim)
# 整段序列:K, V 形状均为 (batch, kv_heads, seq_len, head_dim)
print("per-token KV 元素数 =", 2 * layers * kv_heads * head_dim) # 随层数线性
把「推理是带宽受限」量化:每解一个 token,要把整段 KV Cache(几十 GB)从 HBM 读一遍,而算力只用到几次矩阵乘。算术强度极低,瓶颈是 HBM 带宽而非 FLOPs——这就是 GQA / 量化能直接换成吞吐的根本原因。
python# 解码一步的访存 vs 算力(证明是 memory-bound)
def decode_cost(layers, kv_heads, head_dim, T, d_model, dt=2):
kv_bytes = 2 * layers * kv_heads * head_dim * T * dt # 每步要读的 KV
flops = 2 * layers * (2 * d_model * T) # 粗略注意力 MACs
return kv_bytes, flops
kv, fl = decode_cost(80, 8, 128, 32768, 4096)
print(f"读 KV = {kv/1024**2:7.1f} MB 算力 = {fl/1e9:6.2f} GFLOP")
print(f"算术强度 = {fl/kv:.4f} FLOP/byte (H100 约 300,远低于此)")
# 预期:算术强度不足 1 -> 每步耗时基本等于「读一遍 KV」的时间
# 结论:解码阶段优化的是访存(GQA/量化/分页),不是 FLOPs
7.2 一个具体的显存测算例子
取一个 70B 级别、现代 GQA 配置(约 80 层、64 注意力头、GQA 8 个 KV 头、head_dim 128,近似 Llama 类架构),FP16(2 字节),batch = 1。先算每 token 的 KV 元素:2 × 80 × 8 × 128 = 163,840,约 0.3125 MB(FP16)。再乘以上下文长度,得到单序列显存随长度线性膨胀:
pythonL, g, dh, dt = 80, 8, 128, 2 # 层数, KV头数, head_dim, FP16字节
per_tok_mb = 2 * L * g * dh * dt / 1024**2
for ctx in (4096, 32768, 131072):
gb = kv_cache_gb(L, ctx, g, dh, batch=1, dtype_bytes=dt)
print(f"ctx={ctx:>7} 单序列 KV Cache = {gb:6.2f} GB (每 token {per_tok_mb:.4f} MB)")
# 对照:若用 MHA(kv_heads=64 而非 8),同配置 32K 上下文:
print("MHA 32K =", round(kv_cache_gb(L, 32768, 64, dh, dtype_bytes=dt), 2), "GB")
print("GQA 32K =", round(kv_cache_gb(L, 32768, g, dh, dtype_bytes=dt), 2), "GB")
print("GQA 把 KV Cache 降到 MHA 的", round(g/64, 4), "(即 1/8)")
# 若再叠加 FP8(dt=1),GQA 32K 进一步减半到约 1.3 GB
| 上下文 | 单序列 KV(GQA, FP16) | 含义 |
|---|---|---|
| 4K | ≈ 1.28 GB | 短对话,单卡可轻松承载 |
| 32K | ≈ 10.2 GB | 长文档,已占去一张卡相当一部分显存 |
| 128K | ≈ 40.96 GB | 接近单张 80GB 卡的极限,必须多卡或量化 |
| 1M (GQA, FP8) | ≈ 320 GB | 超长上下文需专家并行 / 量化 / 分页 KV |
python# 部署容量规划:给定显存能放多少并发序列
def max_concurrency(gpu_gb, weights_gb, layers, kv_heads, head_dim,
ctx, dt=2, overhead=0.15):
free = gpu_gb * (1 - overhead) - weights_gb
per_seq = 2 * layers * kv_heads * head_dim * ctx * dt / 1024**3
return int(free / per_seq)
# 80GB 卡、权重 16GB(8B 模型 BF16)、80 层、GQA 8 头、head_dim 128
for ctx in (8192, 32768, 128000):
per = 2 * 80 * 8 * 128 * ctx * 2 / 1024**3
print(f"ctx={ctx:>7} 单序列 KV≈{per:5.2f}GB 可并发≈{max_concurrency(80,16,80,8,128,ctx)} 条")
# 预期:ctx=8K 可并发几十条;ctx=128K 只剩个位数
# 结论:超长上下文与高并发是矛盾的,必须靠量化/GQA/分页来平衡
| 场景 | 上下文 | KV 策略 | 单卡并发 |
|---|---|---|---|
| 短对话客服 | 4K–8K | BF16 GQA | 高(数十条) |
| 长文档问答 | 32K–128K | FP8 GQA + 分页 | 中(个位数) |
| 超长 Agent 轨迹 | 1M | MLA + 量化 + 多卡 | 低(需专家并行) |
- 默写公式:KV 显存 = ? 判据:2×L×n_kv×d_head×T×B×dtype_bytes。
- 手算:L=80、n_kv=8、d_head=128、T=128K、B=1、FP16。参考答案:约 40.96 GB。
- 对比:同配置 MHA(n_kv=64)要多少?参考答案:约 327 GB(GQA 的 8 倍)。
- 量化收益:FP8 量化 KV 能省多少?判据:字节减半(dtype_bytes=1),显存约减半。
- 工程:为什么 PagedAttention 能提吞吐?判据:减少显存碎片、提升可并发放置的序列数。
8. 长上下文工程:切分、外推与遗忘
学习路径
- 读 8.1:理解上下文并行切序列 O(T/N),Ring Attention 用在线 softmax 跨块累加
- 梳理 CP/SP/TP/EP 各自切分轴,说明与 FSDP/TP 的 3D/4D 组合
- 完成本章自测:解释 CP 解决什么、1M 单卡 KV 约 320 GB
核心知识点详解
- 上下文并行切序列 O(T/N):把长度为 T 的序列沿长度切成 N 段给 N 卡,单卡注意力复杂度从 O(T²·d) 降到 O((T/N)²·d),同时 KV 显存每卡 O(T/N)。
- Ring Attention 用在线 softmax 跨块累加:各卡只持有本段的 K/V,环形逐轮转发,用与 FlashAttention 相同的在线 softmax 增量合并结果,不物化完整全局 attention。
- 与 FSDP/TP 的 3D/4D 组合:FSDP 切参数(含通信)、TP 切张量维、SP 切序列维;典型 3D 是 FSDP+TP+SP,4D 再加入 CP 或 EP。每个并行轴各切一份资源,显存靠合力摊薄。
- 1M 单卡 KV 约 320 GB:0.3125MB×1e6 = 320GB,单卡根本放不下,正是 CP 必须存在的原因。常见坑:只加 CP 不加在线 softmax,等价的「全局 softmax 物化」仍会把显存顶爆。
学习路径
核心知识点详解
- Lost in the Middle 现象:长上下文中模型对中间位置的证据召回与利用明显弱于首尾。越长的序列稀释越严重,是 RAG 与重排要解决的问题。
- 三种缓解办法:① RAG 只在上下文里放最相关文本,缩短有效长度;② 关键证据重排到开头/结尾(黄金位);③ 用注意力/位置增强(如延长训练、中间采样)。
- 有效上下文长度 ≠ 标称长度:模型可能「读得进」长文本(不丢单条信息)但「用不好」(跨位置推理弱)。要分别测召回与多跳利用,别只看能塞进去的长度。
- Needle-in-a-Haystack 的局限:只测「单条关键信息能否召回」,不测长程推理与累积,通过 NiH 不代表上下文能力强。常见坑:用 NiH 全通过宣称长上下文好,业务长程任务照样崩。
核心知识点详解
- CP 解决什么:序列太长单卡放不下 KV 与 O(T²) 计算时,用上下文并行把序列切到 N 卡,单卡动态降为 O((T/N)²)。
- 1M 单卡 KV ≈ 320 GB:0.3125MB×1e6 = 312.5GB ≈ 320GB,远超单卡容量,故 1M 上下文必须 CP + 量化 + GQA 合力。
- 三种中段遗忘缓解的收尾:RAG 截短、证据前移、注意力/长训增强,选一个落地即可。常见坑:缓解后不度量「有效上下文」,凭观感就直接上线。
8.1 上下文并行与 Ring Attention
要让单卡装下 128K 甚至 1M token 的注意力,必须把序列维切开。上下文并行(Context Parallelism, CP)沿序列维度把 token 分给 N 张卡,每张卡只持有一段(长度 T/N),显存从 O(T) 降到 O(T/N)。难点在于:第 i 张卡的 query 要计算与「所有卡上的 K/V」的注意力——否则注意力就不完整。
Ring Attention 的解法是把 K/V 块在 GPU 环形拓扑里逐跳传递:每张卡先用本地 K/V 与本地 Q 算一块注意力,再把 K/V 块发给下一台、同时接收上一台传来的块,用「在线 softmax」的 running max/sum 累加跨块的归一化项。通信与计算完全重叠,因此额外通信开销被摊薄。它是 2026 年 1M token 级训练/推理的底层支柱之一。
| 并行方式 | 切分轴 | 与长上下文的关系 |
|---|---|---|
| 上下文并行 CP | 序列 | 直接把长序列摊到多卡,显存线性下降 |
| 序列并行 SP | 序列(激活) | 常与 TP 配合,降低单层激活 |
| 张量并行 TP | 隐藏维 / 头 | 单层放不下时的补充,不直接解决长序列 |
| Ring Attention | 跨卡 K/V 传递 | 让 CP 下的注意力保持完整且通信重叠 |
| 专家并行 EP | 专家 | MoE 专属,与序列长度无关 |
python# Ring Attention 的「在线 softmax」核心:跨块累加而无需全局矩阵
# 每个设备持有一段 query q_local 与一段 kv 块;沿环传递 kv 块
def ring_attn_step(q_local, k_block, v_block, state):
# state: (running_max, running_sum, acc) 在线 softmax 的递推量
s = q_local @ k_block.transpose(-1, -2) / (q_local.shape[-1] ** 0.5)
block_max = s.max(-1, keepdim=True)
m_new = torch.maximum(state.m, block_max)
P = torch.exp(s - m_new) # 本块未归一化权重
acc = state.acc * torch.exp(state.m - m_new) + (P @ v_block)
sum_new = state.sum * torch.exp(state.m - m_new) + P.sum(-1, keepdim=True)
return State(m=m_new, sum=sum_new, acc=acc)
# 实际实现还要叠加因果掩码与跨设备 all-gather/p2p,此处只演示「可分段累加」
CP 的收益与代价要算清:把序列切成 N 段,KV 显存降到约 1/N;但每张卡仍要看全部 K/V,Ring Attention 用环形 p2p 传递 K/V 块、使通信与计算重叠,净开销可压到接近零。代价是 N 越大、环传递次数越多,对互联带宽要求越高。
python# CP 的显存收益与通信量估算
def cp_analysis(L, T, kv_heads, head_dim, N, dt=2, bw_gbps=1800):
kv_full = 2 * L * T * kv_heads * head_dim * dt / 1024**3
traffic = 2 * L * kv_heads * head_dim * T * dt * (N - 1) / N
t_comm = traffic * 8 / (bw_gbps * 1e9) # 秒
print(f"N={N:2d} 每卡KV={kv_full/N:6.2f}GB 通信量={traffic/1024**3:6.2f}GB 约{t_comm*1000:6.2f}ms")
for N in (1, 4, 8, 16):
cp_analysis(80, 1_000_000, 8, 128, N)
# 预期:N 从 1 到 16,每卡 KV 从约 320GB 降到约 20GB;通信量随 N 增大
# 结论:CP 用「多卡 + 环形通信」换取 1M 上下文的可行性
8.2 长上下文的「中间遗忘」现象
一个反直觉的事实:把关键信息放在超长 prompt 的中段,模型表现往往比放在开头或结尾更差——这被称为「Lost in the Middle」(Liu 等,2023–2024)。两条主因:① 注意力在长序列里被稀释,模型对中段 token 分配的权重偏低;② 预训练与 SFT 的监督信号大多集中在文档首尾(如段落首尾、对话头尾),模型「习惯」关注两端。
| 缓解手段 | 原理 | 代价 |
|---|---|---|
| 检索增强(RAG) | 只把相关片段送入上下文,缩短有效长度 | 需检索质量高,否则引入噪声 |
| 重排序 / 重打包 | 把关键证据移到首尾 | 需先识别关键片段 |
| 长上下文续训 | 用长样本继续预训练,弥合位置外推 | 算力最贵但最可靠 |
| 压缩 / 摘要中段 | 先摘要再拼回 | 可能丢细节 |
| 位置插值 + 温度缩放 | 缓解 RoPE 外推衰减 | 过长仍有损 |
python# 「大海捞针」位置敏感性测试:把关键信息放在不同深度看命中率
def needle_test_positions(depths=(0.0, 0.25, 0.5, 0.75, 1.0)):
for d in depths:
# probe(model, long_doc_with_needle(d)) -> 是否答对
print(f"信息位于相对位置 {d:4.2f} -> 记录命中率")
# 典型结论:首尾命中率高、中段明显下降(Lost in the Middle)
needle_test_positions()
# 预期:d=0.0/1.0 命中率接近满分,d≈0.5 明显下降
# 对策:关键证据前置/后置、RAG 缩减有效长度、长上下文续训
| 评测(2026) | 测什么 | 容易被高估的点 |
|---|---|---|
| Needle-in-a-Haystack | 单点检索 | 只考「找到」,不考「推理」 |
| RULER / LongBench | 多任务长上下文 | 仍需按领域实测 |
| 有效上下文长度 | 能答对的最长长度 | 标称窗口常远大于有效窗口 |
| 跨段一致性 | 多段信息整合推理 | 中段证据易被忽略 |
- 解释 CP:上下文并行解决什么问题?判据:把序列摊到多卡,显存从 O(T) 降到 O(T/N)。
- Ring 机制:为什么 Ring Attention 能保持注意力完整?判据:K/V 块沿环传递 + 在线 softmax 累加。
- 算收益:1M 上下文、80 层、GQA 8 头、head_dim 128,单卡 KV 约多少?参考答案:约 320 GB,故必须多卡 CP / 分页 / 量化。
- 诊断:模型「支持 128K」但长文档中段问题答错,最可能原因?判据:Lost in the Middle + 外推衰减,非单纯显存问题。
- 对策:给出三种缓解中段遗忘的办法。判据:RAG 缩长、关键证据重排到首尾、长上下文续训 + 位置缩放与温度。
项目里程碑
Hamauls Orion 的模型层自己实现:手写 Multi-Head Attention、RoPE、RMSNorm、SwiGLU、KV Cache,并做到能加载开源权重、输出与 HuggingFace 参考实现数值一致。这一步打通「架构 ↔ 权重 ↔ 推理」的完整链条。
本阶段产出(直接进入项目仓库)hamauls_orion/models/transformer.py:自研 GPT(RoPE / RMSNorm / SwiGLU / GQA / KV Cache)tests/test_parity.py:与 HuggingFace 同权重逐层 logits 对比,max abs diff < 1e-3hamauls_orion/models/kv_cache.py:增量解码实现 + 显存换算脚本,与 M2 的公式互相验证docs/model/architecture.md:架构图 + 参数量/FLOPs/KV 显存逐项拆解- 投机解码最小实现(draft + verify),记录加速比
阶段练习项目
- 在字符级 / 小词表语料上从零训到 PPL 相对随机初始化明显下降(如 <2.0),生成出一到两句通顺片段
- 训练全程可复现:固定 seed 后同参数两次训练,loss / PPL 曲线一致到 ±2%
- 产出 Top-K / 温度采样对比并用温度 τ=0(贪心)与 τ>0(采样)各生成样例验证差异
- 手写 Token / Pos Embedding、RMSNorm+Pre-LN、GQA 注意力、SwiGLU FFN,禁用
nn.TransformerEncoder与nn.Transformer - 实现带因果掩码的前向与
F.cross_entropy(logits, labels)训练循环,逐位置预测下一 token 并记录每步 PPL - 生成阶段实现 top-k + 温度采样、生成时裁剪上下文长度,支持 batch 解码
- 用一个单元测试验证因果掩码已置未来为 -1e9(权重不被未来令牌污染)
- 记录 embedding dim 与 vocab 的形状不依赖硬编码,验证 weight tying 是否开启
minigpt.py(可训练可生成)+ppl_curve.png- Top-K / 温度采样对比表与实现说明 md
不做真实大规模语料与多卡并行训练,仅做字符级 / 小词表的能力验证。
- 在 4K / 32K 下测出 MHA / GQA / MQA 三种注意力的 KV 显存与生成吞吐,形成可复现对比表
- 实测 KV 显存与公式 2·L·kv_heads·head_dim·seq·dtype 计算值偏差 <10%
- 给出部署建议,并说明为何 memory-bound 下 GQA 比 MHA 更能提吞吐与并发
- 实现 MHA(g=h)、GQA(g
- 用
kv_cache_bytes/ 显存探测工具分别记录 4K 与 32K 上下文的缓存大小,与公式互验- 测解码吞吐(tokens/s),对比三种方案在同一模型下的差异,并注明 OOM 边界
- MHA 相对 GQA 的 KV 显存应约为 g/h 倍(核验其 8 倍关系之一)
- 用
kv_compare.py+ 对比表(4K / 32K × 三方案)数据- 部署建议 md(含 GQA 组数选择与并发容量规划)
不做 FP8 / INT4 量化的 KV 实现,仅评估 bf16 基线下的相对差异。
- 在固定 4K 训练基座上,于 8K 评测集比对四种外推方案,得到可复现的 PPL 差异
- 标出现实有效窗口与标称窗口的差距,并解释直接外推为何退化最快
- 任一方案的最终 PPL 都应优于直接外推(naive)数个百分点才算有效
- 训练一个窗口 4K 的 GPT 作为固定 base,冻结后仅改位置编码与注意力温度做外推
- 实现直接外推 / PI / NTK-aware / YaRN 四种,仅在 RoPE base 与温度 τ 上做改动
- 在 8K 长序列评测集上分别算 PPL,并记录 4K 内(外推未超限)作为参照基线
- 统一 dtype 与 attention 内核(math vs flash),避免 kernel 差异污染对比
- 记录各方案的有效窗口,说明其与标称窗口不一致的原因
extrap_experiment.py+ 四方案 PPL 对比表与曲线- 结论 md(含有效窗口测量与选型建议)
不做长上下文续训 / 退火微调优化外推,仅评估固定基座下外推方法的相对表现。
- 产出「架构 / 数据 / 后训练 / 基建 / 评测」五栏结构清晰的精读笔记
- 逐条标出与 2017 原版 Transformer 的差异并给出动机,条数不少于 8 条
- 笔记能直接支撑选型与后续手写实现决策(含关键数值与引用)
- 选定 DeepSeek-V4 或 Qwen3.5 并通读其技术报告
- 五栏各提取关键数字(参数量 / 层数 / 词表 / RoPE base / 数据量 / 超参)并标注来源引用
- 单列一栏逐条比对与 2017 原版的差异(如 Pre-LN/RMSNorm/GQA/RoPE/SwiGLU/MLA/MoE)并给理由
- 在「评测」栏给出作者自评与第三方独立评估的差距对照
report_notes.md结构化精读笔记
不做复现训练,只做文献精读与要点提取。
- 无需
nn.TransformerEncoder手写全部组件,在小语料上训到能续写通顺片段 - 复现选定开源权重:固定输入下 final logits 的
max abs diff < 1e-3、cos > 0.9999 - tokenizer / RoPE base / RMSNorm eps / 权重共享四要素对齐并在 architecture.md 文档化
- 手写 embedding(含可选 weight tying)、RMSNorm+Pre-LN、GQA、RoPE、SwiGLU、MoE(按选择),逐层参数量与权重文件对得上
- 实现增量解码 kv_cache,生成功能可用
- 加载权重后用逐层输出 cos 对比定位第一个偏差层,迭代收敛到 final logits 判据
- 以
pytest输出test_parity.py,判定max abs diff < 1e-3与cos > 0.9999 - 在 architecture.md 记录参数/FLOPs 估算与四要素配置值,供与参考实现核对
- 可加载开源权重的
gpt.py+kv_cache.py+test_parity.py architecture.md(结构 / 参数量 / FLOPs / 对齐要点)
不做多轮 / 指令能力级预训练,不对齐闭源权重(只选公开权重),不做推理部署与量化。
常见误区
- 只记结构不记复杂度,面试被问「为什么长上下文贵」答不上来。
- 以为 FlashAttention 是近似算法,或者以为它减少了 FLOPs(其实减少的是显存搬运)。
- 把 MoE 当成「免费扩容」,忽略显存驻留与通信代价,部署时才发现放不下。
- 死记 MHA / MQA / GQA 的定义,却算不清 KV Cache 显存,无法做容量规划。
- 忽略位置编码,导致微调或推理时长上下文表现崩坏却找不到原因。
- 只看排行榜选模型,不做自己的小规模评测,上线后才发现能力维度不匹配。
- 把预训练当成唯一手段,忽视数据质量与后训练带来的巨大收益。
面试高频问题速答
写出注意力的公式并解释每个符号。
Attention(Q,K,V) = softmax(QKᵀ/√d_k)V。Q 是查询矩阵(当前 token 想找什么),K 是键矩阵(每个位置能提供什么),V 是值矩阵(每个位置的实际内容),d_k 是每个头的维度,√d_k 用于稳定数值。复杂度 O(T²d):T² 来自 QKᵀ,乘以 d 是每个点积的代价。
为什么 2026 年主流模型都改用 GQA?
推理阶段是显存带宽受限的,KV Cache 是最主要的显存占用方(随上下文长度线性增长)。GQA 让多个查询头共享一组 KV 头,在几乎不损失质量的前提下把 KV Cache 缩小数倍(例如 32 头 KV → 8 头,显存降到 1/4),从而支持更长上下文与更大并发。MQA 更进一步只保留 1 组 KV,节省最多但质量损失更明显。
Pre-LN 和 Post-LN 有什么区别?
Post-LN 把归一化放在残差相加之后(原版),深层时梯度需穿过多次 LN,容易需要精细的 warmup 才能训稳。Pre-LN 把归一化放在子层输入之前,残差路径上没有归一化,形成一条干净的恒等通路,梯度更稳、对 warmup 不敏感,因此成为现代默认。
MoE 的负载均衡问题是什么?怎么解决?
router 训练中容易出现「赢者通吃」:热门专家被更多 token 选中 → 梯度更多 → 更强 → 更热,最终少数专家承担几乎所有计算,其余专家完全没被训练,等效容量大幅缩水。解决手段:辅助负载均衡损失(鼓励分配比例与路由概率都均匀)、专家容量上限与 token 丢弃、无辅助损失的偏置调整(DeepSeek 的做法)、以及对 router logits 做噪声扰动。
Scaling Law 告诉你给定算力该怎么分配参数与数据?
Chinchilla 的结论是参数与训练 token 应大致同比例增长,最优比例约 20 tokens/参数。但 2026 年的生产实践普遍「刻意过训练」——用 100+ tokens/参数训练较小的模型,因为推理成本是长期主导成本,小模型的单位推理成本低得多,且便于部署。此外数据质量、去重与配比带来的收益常常超过单纯扩大规模。
怎么把 32K 上下文扩展到 128K?
三条路:① 位置编码缩放(线性插值、NTK-aware、YaRN),改动小、成本低,但过长时质量衰减;② 稀疏/滑动窗口注意力降低计算量,但可能丢远距离依赖;③ 直接用长样本继续预训练(最可靠也最贵)。实践中常组合使用:先做位置插值让模型「能跑」,再用少量长样本文本做续训让模型「跑得好」,并配合注意力温度缩放。