← 返回学习路线 ◆ 贯穿项目
大模型核心 · 阶段 8 · Transformer 与现代大模型架构
Stage 08 / 17 · 大模型核心

Transformer 与现代大模型架构 Transformer & Modern LLM Architecture

这是整条路线的重心。2026 年的架构演化几乎全部发生在这一层:注意力的计算方式(FlashAttention → 稀疏 / 线性)、参数的组织方式(稠密 → 细粒度 MoE)、位置编码的外推能力(RoPE → 各种缩放变体)、优化的数值特性(AdamW → Muon)。把这些「为什么」讲清楚,你就具备了读懂任何新模型技术报告的能力。

⏱ 5–6 周 🎯 进阶 · 重中之重 ◆ 里程碑 M8 2026-09-29
TransformerAttentionRoPEMoEGQAFlashAttention长上下文

阶段总览

✔
学完你能做到
  • 能手推 Scaled Dot-Product Attention 与多头注意力,说清每一步的维度与复杂度
  • 能从零实现一个可训练的 MiniGPT,并解释每个模块存在的理由
  • 掌握 RoPE / ALiBi 等位置编码,理解长上下文外推的原理与手段
  • 理解 MHA / MQA / GQA 的取舍,能算清 KV Cache 的显存占用
  • 理解 MoE 的路由、负载均衡与专家并行,能解释为什么 2026 年主流都是 MoE
  • 理解 FlashAttention 为什么快(IO 感知,不是近似),以及线性注意力 / Mamba 的动机与代价
  • 能读懂一份主流模型的技术报告,并做出选型判断
阶段知识结构总览 · 讲透注意力,看懂 2026 架构选择
阶段 8 · Transformer 与现代大模型架构Transformer & Modern LLM Architecture · 8 大章 · 100+ 知识点
1. Attention 起点Scaled Dot-Product多头的算力账本MHA / MQA / GQA / MLAKV Cache 显存公式QK-Norm
2. 现代架构改进Pre-LN + RMSNorm + SwiGLU细粒度 MoE + 共享专家RoPE 外推 YaRNFlashAttention IO 感知线性 / 混合注意力
3. 主流模型横评六家前沿模型格局按负载加权打分LiteLLM 统一网关降级链 + 影子流量
4. 从零实现 MiniGPTDecoder-only 结构参数量拆解权重共享 weight tyingstd=0.02 初始化对齐开源权重 cos>0.9999
5. Scaling LawKaplan → Chinchilla过训练 100+ tokens/paramC ≈ 6NDMFU 30–50%DP / TP / PP / EP / ZeRO
6. FFN / 归一化 / 残差FFN 占 2/3 参数SwiGLU 8/3 dPre-LN 梯度直通RMSNorm 省均值残差缩放 ReZeroMuon
7. KV Cache 推理工程KV 显存公式FP8 / INT4 KV 量化PagedAttentionmemory-bound 算术强度并发容量规划
8. 长上下文工程上下文并行 CPRing Attention 在线 softmaxLost in the MiddleRAG / 重排缓解有效上下文长度
贯穿项目 · M8 从零实现 GPT 并对齐开源权重第 39–46 周自研 GPT(RoPE / RMSNorm / SwiGLU / GQA / KV …与 HuggingFace 同权重逐层 logits 对比,max abs diff…增量解码实现 + 显存换算脚本,与 M2 的公式互相验证架构图 + 参数量/FLOPs/KV 显存逐项拆解
学习路径
★
学习策略:一条主线 + 三条支线:主线是注意力本身(本质、复杂度、变体)。支线一是「怎么让训练更稳」(归一化、激活、初始化、优化器)。支线二是「怎么让推理更省」(KV Cache、GQA、FlashAttention、量化)。支线三是「怎么让上下文更长」(RoPE 外推、稀疏与线性注意力)。2026 年的所有架构新闻,基本都能归到这三条支线上。
周次主题交付物
第 1 周Attention 本质与复杂度手推 + numpy 实现 + 复杂度分析
第 2 周Transformer 全结构从零实现可训练 MiniGPT
第 3 周现代组件与位置编码RoPE / RMSNorm / SwiGLU 手写 + 外推实验
第 4 周MoE 与高效注意力阅读 DeepSeek / Qwen 技术报告并写笔记
第 5 周长上下文与推理成本KV Cache 显存测算 + GQA 对比实验
第 6 周Scaling Law 与模型横评选型决策文档

1. Attention:一切的起点

知识结构图 · Attention 起点
Attention 起点3 大知识域 · 16 个知识点
注意力本质Q / K / V 检索模型softmax(QKᵀ/√d_k)V÷√d_k 防 softmax 饱和因果掩码O(T²·d) 复杂度温度 τ 缩放
学习路径
  1. 读 1.1:默写 softmax(QKᵀ/√d_k)V,背下 ÷√d_k 防 softmax 饱和与 O(T²·d)
  2. 跑 numpy 实现验证每行权重和为 1、因果掩码让第一行只关注自己
  3. 实测序列长度翻倍耗时约 4–16 倍,坐实 T² 复杂度
  4. 对接 M8:手写 GPT 时实现并校验因果注意力的数值正确性
✔ 能手写并测试 ScaledDotProductAttention,解释缩放因子与掩码作用
核心知识点详解
  • 缩放 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 外推里调注意力温度,本质就是调这个 τ 以补偿未见过的长距离。
多头与复杂度账本4d² 投影算力注意力算力与头数无关KV Cache 是推理显存大头QK-Norm 防熵坍缩SDPA 自动选内核
学习路径
  1. 读 1.2:背下 QKV 合并投影 4d² 算力、注意力算力与头数无关、KV Cache 显存公式
  2. 跑 MultiHeadAttention 前向,验证 SDPA 自动选内核与输出 shape
  3. 用 attn_flops 函数验证 h=8/16/32/64 时注意力 FLOPs 不变
  4. 对接 M8:手写 MHA 并与 PyTorch SDPA 对齐输出,说明 QK-Norm 的作用
✔ 能手写 MHA,实测注意力算力与头数无关,输出与参考实现数值一致
核心知识点详解
  • 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 开关导致数值对不上。
MHA / MQA / GQA / MLAGQA KV 降 g/hMHA / GQA / MQA 取舍MLA 低维潜向量 c_kvmemory-bound 推理Qwen3 统一 8 KV 头
学习路径
  1. 读 1.3:区分 MHA / MQA / GQA,理解 KV Cache 正比于 KV 头数、memory-bound 推理
  2. 用 kv_cache_bytes 对比 MHA/GQA/MQA 在 32K 上下文下的显存(约 64/16/2 GB)
  3. 读 MLA 的低维潜向量 c_kv 方案,对比其与 GQA 的 KV 量级
  4. 对接 M8:手写 GPT 时选用 GQA,权衡推理显存与吞吐
✔ 能算清 GQA 相对 MHA 的 KV 显存下降 g/h,并解释为何 2026 主流用 GQA/MLA
核心知识点详解
  • 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))          # 因果掩码的效果
⚠
为什么要除以 √d_k:若 Q、K 各维独立且方差为 1,点积的方差约为 d_k。d_k 很大时(如 128)点积数值范围很宽,softmax 会被推入饱和区——梯度趋近于 0,训练停滞。除以 √d_k 把方差拉回 1,这就是「缩放」的全部理由,也是非常经典的面试题。

把 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 较小时反而是这部分主导
FFNO(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最大最佳最低小模型 / 质量优先
GQAg (1中等(h/g 倍缩小)接近 MHA高2026 主流大模型默认
MQA1最小略损最高超低延迟 / 流式场景

为什么 2026 年主流模型几乎全用 GQA?根因是推理是显存带宽受限(memory-bound)而非算力受限:每解码一个 token,都要把整段 KV Cache 从 HBM 读进计算单元。Cache 越小 → 单卡能塞的 batch 越大、访存越少 → 吞吐越高、延迟越低。GQA 在几乎不掉点的前提下把 Cache 砍掉 4–8 倍,是「更长上下文 + 更高并发」与「质量」之间的最优折中。MQA 更激进但质量损失更明显,只在极致的延迟敏感场景(如部分语音/流式)出现。

★
选型记忆点:KV 头数 = 显存与吞吐的旋钮。需要超长上下文或高并发 → 压低 KV 头数(GQA);质量绝对优先且不愁显存 → MHA;DeepSeek 的 MLA 则另辟蹊径,用低秩投影把 KV 压到极小维度,等效于「更极致的 GQA」,是 2026 另一主流路线。

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,V1×早期 LLaMA
GQAg 组 K,Vg/hLlama 4 / Qwen3 / Mistral
MQA单组 K,V1/h部分流式 / 语音模型
MLA低维潜向量 c_kv远小于 GQADeepSeek-V3/V4 / Kimi
★
1.4 动手练习与自测:本节是全站地基,题目要能「手推 + 手算」才算过关。
  1. 默写公式:写出 Scaled Dot-Product Attention,标注 Q/K/V 形状与复杂度。判据:softmax(QKᵀ/√d_k)·V,复杂度 O(T²·d)。
  2. 解释缩放:为什么除以 √d_k 而不是 d_k?判据:点积方差≈d_k,除 √d_k 把方差拉回 1,避免 softmax 饱和、梯度消失。
  3. 改代码:把 1.1 的 numpy 实现去掉因果掩码,观察第一行权重。判据:无掩码时第一行不再集中在自己、权重摊平。
  4. 手算 KV Cache:80 层、GQA 8 头、head_dim 128、FP16、32K、batch=8。参考答案:2×80×8×128×32768×8×2 B ≈ 81.9 GB。
  5. 选型判断:要支持 1M 上下文且高并发,KV 头数往哪调?判据:压低(GQA/MQA/MLA),因为 KV Cache 正比于 KV 头数。

2. 现代架构改进(2026 视角)

知识结构图 · 现代架构改进
现代架构改进4 大知识域 · 24 个知识点
现代 Decoder 块Post-LN → Pre-LNLayerNorm → RMSNormReLU → SwiGLU绝对编码 → RoPE块参数 ≈ 12d²
学习路径
  1. 读 2.1:用「对照清单」读透一个现代 Decoder 块(Pre-LN/RMSNorm/SwiGLU/RoPE/GQA/FlashAttention)
  2. 搭建一个现代 Decoder 块并将各位置与 2017 原版逐项对比
  3. 对接 M8:手写 GPT 时按此清单组织 block,逐层对齐结构与参数量
✔ 能逐项交代现代块相对原版的十处改动及其动机,叠出 12d² 级参数块
核心知识点详解
  • 现代块 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 差出一个量级。
MoE 参数组织Top-K 路由 + 共享专家负载均衡 L_aux=α·Σf_i·P_iDeepSeek 无辅助损失偏置router z-loss专家容量因子 1.25–2.0专家并行 EP all-to-all激活比 3–4%
学习路径
  1. 读 2.2:理解 Top-K 路由、共享专家、负载均衡与激活比 3–4% 的量级
  2. 跑一个 MoE 块计算参数量与激活参数,验证"总参大但激活少"
  3. 对接 M8:在 architecture.md 里用 FFN 改造的视角说明 MoE 与标准 FFN 的参数量差
✔ 能算出 MoE 的总参数 vs 激活参数,并解释"大参数、低激活、高性价比"
核心知识点详解
  • 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 送到对应专家所在卡,通信随专家数上升。常见坑:只算「总参大」就以为省显存,忘了激活参数与各专家权重都要驻留显存。
位置编码与长上下文外推RoPE 相对位置旋转位置插值 PINTK-aware 缩放YaRN 温度补偿ALiBi 线性偏置有效窗口 ≠ 标称窗口
学习路径
  1. 读 2.3:理解 RoPE 相对位置旋转、PI / NTK / YaRN 外推,并知道有效窗口 ≠ 标称窗口
  2. 实现 RoPE 前向并测其旋转不变量,理解相对位置性质
  3. 跑一次长上下文外推,记录标称与有效窗口的差距
  4. 对接 M8:手写 GPT 的 RoPE 并按目标 base 复现,与参考实现对齐
✔ 能实现 RoPE 并复现其相对位置性质,定量说明外推小于标称窗口的原因
核心知识点详解
  • 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+,长上下文表现直接崩。
FlashAttention 革命IO 感知分块在线 softmax显存 O(T²) → O(T)FA1 / FA2 / FA3 / FA4 差异滑动窗口 / 线性注意力Mamba / SSM 状态空间
学习路径
  1. 读 2.4:理解 FlashAttention 的 IO 感知分块与在线 softmax,显存 O(T²)→O(T)
  2. 对比朴素注意力与 SDPA 的显存/速度,测量长序列下的差异
  3. 完成自测:说明 RA1–FA4 迭代与滑动窗口/线性注意力的取舍
  4. 对接 M8:在 architecture.md 说明注意力实现为何用 SDPA/FlashAttention
✔ 能量化 FlashAttention 省下的显存并解释在线 softmax,说明 Mamba 的另一路线
核心知识点详解
  • 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-LNPre-LN(归一化在子层之前)Pre-LN 梯度更稳,不需要精细的学习率预热
归一化类型LayerNormRMSNorm省掉均值计算,更快,效果相当
激活函数ReLUSwiGLU / GeGLU门控带来更好的表达力,代价是参数量增加
位置编码正余弦绝对编码RoPE(含各种缩放变体)相对位置、可外推到更长上下文
注意力头MHAGQA / MQA大幅削减 KV Cache
注意力实现朴素 matmul + softmaxFlashAttention / SDPAIO 感知的分块计算,显存线性、速度更快
FFN 组织稠密 FFN细粒度 MoE(共享专家 + 路由专家)总参数量大但激活参数少,性价比高
优化器AdamAdamW → Muon解耦权重衰减;Muon 做谱范数正交化提升稳定性
注意力架构纯二次注意力混合:滑动窗口 / 稀疏 / 线性注意力 + 少量全局层把 O(T²) 压到接近 O(T)
训练精度FP32BF16 / 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 级别)。

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
★
MoE 的真实代价:MoE 省的是计算,不省显存——所有专家参数都得驻留显存,所以部署 MoE 需要大量显存或专家并行 + offload。另外它带来训练不稳定性(路由抖动、负载不均)与推理时的通信开销。理解这些代价,你才能回答「为什么某些场景稠密模型反而更合适」。

负载均衡的数学形式:设 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 / 高速互联依赖极强。

维度稠密 DenseMoE(如 8 专家激活 2)
总参数N数倍于 N(容量大)
每 token 激活参数N≈ N/4(省计算)
每 token FLOPs高低(约 1/4)
显存占用N≥ N(全部专家常驻)
训练通信DP/TP/PP额外 EP all-to-all
质量/美元中等更优(同成本更高容量)
ℹ
为什么 2026 前沿几乎都是 MoE:关键不等式:质量随「总参数」提升,而推理成本由「激活参数」决定。把总参数堆大、激活参数保持小,是「既要知识容量又要低单 token 成本」的唯一可扩展路线。DeepSeek-V4(约 1.6T 总参 / 49B 激活)、Qwen3.5(约 397B / 17B)、Kimi K3(约 2.8T)都走此路;细粒度专家 + 共享专家进一步放大组合多样性。只有当显存或互联成为硬约束、且任务不需要超大容量时,稠密模型才更划算。

把「省计算不省显存」量化:总参数 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~49BMLA细粒度 + 共享专家,无辅助损失均衡
Qwen3.5~397B~17BGQA 8 头多语言 + 门控线性注意力混合
Llama 4Maverick/Scout部分激活GQA原生多模态,开源权重
Kimi K3~2.8T~60BMLA/KDA超长上下文 + Agentic
GLM-5.xMoE—GQA/DSA中文与长程 Agent
ℹ
2026 训练侧三个数字:① 负载均衡 aux 系数 α 常取 0.01 量级;② 专家容量因子 capacity_factor 常取 1.25–2.0;③ 训练用 router z-loss + 专家并行 EP,EP 的 all-to-all 是吞吐瓶颈——这是 MoE 集群必须配 NVLink / 高速互联的根因。

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改动小,适合推理端临时扩展
YaRNYaRN: 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 系列32K128KNTK / Dynamic NTK + 续训
DeepSeek-V3/V4128K128K+YaRN + 长样本文本续训
Llama 4256K10M 级位置缩放 + 长上下文训练
Gemini 3 Pro—1M–2M档案级长训练 + 原生多模态
Claude 4.5—1M长上下文训练 + 检索友好
✔
外推的实战口径:「宣称窗口」≠「有效窗口」。经验上,不做续训的纯外推(PI/NTK)在 2–4 倍长度内可接受,超过就明显掉点;要真正用好 8 倍以上长度,必须有长样本文本续训 + YaRN 温度缩放。评测时用「有效上下文长度」(长文档问答仍能答对的最长长度),而不是标称值。

2.4 为什么 FlashAttention 是真正的革命

关键认知:FlashAttention 不是近似算法,它算出的结果与朴素注意力完全一致。它快的原因不是少算了,而是少搬运了——GPU 显存层次中,HBM(大显存)带宽远低于 SRAM(片上缓存),而朴素实现需要把 T×T 的注意力矩阵写回 HBM 再读出来。

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 推到极致。关键不变式:三者算出的结果与朴素注意力逐位一致——它们是精确算法,不是近似。

版本核心改进相对提速硬件
FA1IO 感知分块 + 在线 softmax基线Ampere 及更早
FA2更优 warp 并行 + 去冗余 rescale~2× over FA1Ampere / Ada
FA3异步 GEMM/softmax 重叠 + FP8~1.5–2× over FA2Hopper (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)在 ≤128K、且需要精确长程检索时仍是默认;超过这个量级、且任务允许近似(摘要、长文档分类、流式),才值得上滑动窗口 / 线性 / SSM。混合架构的调参(全局层比例、窗口大小、状态维度)本身就是 2026 年的前沿工程活。

为什么 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 的任务靠那几层全局注意力兜底。

★
2.7 动手练习与自测:现代架构不是「哪个新用哪个」,而是「每个改动各修一个具体问题」。
  1. 对照表默写:说出 2017 原版到 2026 主流在「归一化位置/类型、激活、位置编码、注意力头、实现」五处的变化及理由。判据:Pre-LN、RMSNorm、SwiGLU、RoPE、GQA、FlashAttention 各能说一句 why。
  2. 算激活比:某 MoE 总参 600B、激活 30B,激活比多少?参考答案:5%;每 token 训练 FLOPs 正比于激活参数量。
  3. MoE 失效诊断:训练中某专家吃掉 80% token 怎么排查?判据:看 f_i·P_i 与 router 熵,加 aux loss / 偏置调整 / 调 capacity_factor。
  4. RoPE 外推:解释 PI 与 NTK-aware 的差别。判据:PI 压缩位置(改输入),NTK 增大 base(改频率),YaRN 再叠温度缩放。
  5. FlashAttention 判断:它省的是 FLOPs 还是访存?判据:访存(IO);结果与朴素注意力逐位一致,不是近似。
  6. 混合架构取舍:什么时候不敢用线性注意力?判据:需要精确长程检索(大海捞针、跨段推理)时全局注意力更稳。

3. 主流模型架构横评与选型

知识结构图 · 主流模型横评与选型
主流模型横评与选型3 大知识域 · 13 个知识点
2026 模型格局GPT / Claude / GeminiDeepSeek MLA + MoE + MuonQwen MoE + 线性注意力Kimi 超长上下文开源逼近闭源
学习路径
  1. 读 3.1:梳理 2026 六家前沿与开源阵营,记 DeepSeek MLA+MoE / Qwen MoE / Kimi 长上下文
  2. 对比开源 vs 闭源在编码/数学/Agentic/多模态上的差距分布
  3. 对接 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 元数据的权重,对齐时无从下手。
选型决策清单先定硬约束按负载加权打分LiteLLM 统一网关降级链 fallback影子流量灰度真实数据小规模评测
学习路径
  1. 读 3.2:背下选型决策清单(先定硬约束→加权打分→LiteLLM 网关→降级链→影子灰度→小规模评测)
  2. 给定一个业务约束,按清单选型并解释加权依据
  3. 对接 M8:确认所选权重支持你的 RoPE base / 词表,避免实现与选型冲突
✔ 能按清单给约束下的模型打分排序,并说明软硬约束的取舍
核心知识点详解
  • 先定硬约束再加权打分:硬约束如延迟上限、每千次调用成本、是否允许数据出境、是否必须自托管、合规。先做 constraint_filter 排除,再对剩下的能力维度(长上下文/工具调用/结构化输出/多语言/多模态)加权打分排序,避免「很优但根本不能用」。
  • 用真实数据做小规模评测,别靠排行榜:取 100–300 条自己的真实样本先跑一遍,而不是直接用公开榜单。排行榜三条失败模式:过度拟合榜单题、测试集污染、与你的任务分布不匹配。
  • 算总拥有成本 TCO:API 单价 × 预估 token 量,或自托管 GPU 小时成本 + 运维人力 + 工程时间。开源省的是单价,不是总成本。
  • 用网关隔离,留降级链:用 LiteLLM / 统一网关抽象 provider,配 fallback 与影子灰度,避免供应商锁定。常见坑:只算单价不算长尾与切换成本,结果被单一模型绑架。
动手自测要点日 500 万次短分类选型算月成本三条排行榜失败模式
学习路径
  1. 完成本章自测:日 500 万次短分类的选型与月成本估算,指出三条排行榜失败模式
✔ 自测命中判据,能定量给出选型与成本并识别排行榜陷阱
核心知识点详解
  • 日 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 ClaudeClaude Sonnet 4.5 / Opus 4.7闭源,长上下文 1M,工具调用稳定性强编码、长文档、MCP 生产 Agent
Google GeminiGemini 3 Pro / 3.5原生多模态,2M 上下文视频音频、科研、Workspace 集成
DeepSeekV3 / V4MLA + MoE + Muon,GRPO 后训练自托管前沿、成本敏感场景
QwenQwen3 / 3.5(235B / 397B MoE)MoE + 门控 DeltaNet 线性注意力,多语言强多语言部署、Apache 许可要求
KimiK2 / K3MoE + MLA/KDA,超长上下文长文档、Agentic 编码
GLMGLM-4.5 / 5 / 5.2MoE + DSA,GRPO 起步后长程回归 Critic中文场景、Agent 长程任务
Llama / GemmaLlama 4 / Gemma 3-4开源友好许可,端侧与自托管私有化部署、微调底座
MiniMax / 混元 / MiMoM2.5 / 混元 / MiMo开源权重,编码与 Agentic 表现突出性价比自托管、垂直微调
ℹ
一个实用的判断:选型不要看排行榜总排名,要看你的负载最吃哪一项:需要极长文档 → 看上下文与检索质量;需要多步工具调用 → 看 Agentic benchmark 与工具调用稳定性;需要自托管控成本 → 看开放权重模型的每 token 成本与集群可行性。同一档能力下,专用微调模型的性价比通常高于高两名的通用模型。
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
多步工具调用 AgentAgentic 稳定性 + 工具调用GPT-5 / Claude 4.5 / GLM-5.x
大批量分类抽取成本 + 一致性DeepSeek-V4 / Qwen3.5
私有化 / 数据不出境可自托管 + 许可Llama 4 / Qwen3.5 / Gemma
端侧 / 低延迟小模型 + 量化Qwen3 小杯 / Gemma 3-4

3.2 选型决策清单

先定约束
延迟上限、每千次调用成本上限、是否允许数据出境、是否必须自托管、合规要求。这些是硬约束,先排除不满足的。
再定能力维度
主要是长上下文、工具调用、结构化输出、多语言、多模态中的哪几项?按权重排序。
用真实数据做小规模评测
不要用公开排行榜决定。取 100–300 条自己的真实样本,跑一遍再做决定。
算总拥有成本
API 单价 × 预估 token 量,或自托管的 GPU 小时成本 + 运维人力 + 工程时间。开源省的是单价,不是总成本。
留好替换能力
用 provider 抽象层(如 LiteLLM / 统一网关)隔离模型调用,避免被单一供应商锁定;同时准备好 fallback 与降级策略。
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 可按模型核对预算
# 关键:把「模型选择」收敛为可观测、可回滚的一处配置
★
3.3 动手练习与自测:选型题考的是「把主观判断拆成可复算的数」。
  1. 做决策:给「每天 500 万次短分类」选一个模型并写出理由。判据:以成本与一致性为先,选低价 MoE(如 DeepSeek-V4 / Qwen3.5),而非旗舰。
  2. 算总成本:某 API 每百万 token 0.28 USD,日耗 4 亿 token,月成本?参考答案:约 0.28×400×30 ≈ 3360 USD。
  3. 风险点:列出三条「只看排行榜选型」的失败模式。判据:能力维度错配、成本失控、供应商锁定。
  4. 工程抽象:如何做到换模型不改业务代码?判据:统一网关 / provider 抽象层(LiteLLM),模型名进配置。
  5. 进阶:设计一个「影子流量」灰度方案。判据:小比例真实请求并行跑候选模型,离线比对质量与成本后再切主流量。

4. 从零实现 MiniGPT

知识结构图 · 从零实现 MiniGPT
从零实现 MiniGPT3 大知识域 · 15 个知识点
MiniGPT 结构实现Token + Pos EmbeddingRMSNorm + Pre-LNGQA + SwiGLU交叉熵预训练目标Top-K / 温度采样生成上下文裁剪
学习路径
  1. 读 4.1:逐行理解 MiniGPT 结构(Token/Pos Embedding、RMSNorm+Pre-LN、GQA、交叉熵、Top-K/温度采样)
  2. 亲手实现或敲一遍 MiniGPT,回答"删除某行会怎样"并测生成
  3. 对接 M8:以该实现为基线逐层扩展出可加载开源权重的 GPT
✔ 能独立实现 MiniGPT 并跑通训练与温度采样生成,理解每一行的作用
核心知识点详解
  • 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) 第一个参数是词表大小。
参数量与初始化embedding + 每块 ≈ 12d²权重共享 weight tyingstd=0.02 初始化输出投影零初始化残差缩放 1/√N
学习路径
  1. 读 4.1:背下 embedding + 每块 ≈ 12d²、权重共享、std=0.02、输出投影零初始化
  2. 用代码核对参数量估算,验证 weight tying 与残差缩放 1/√N
  3. 对接 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 实现是否一致,再谈训练稳定。
对齐开源权重复现 tokenizer复现 RoPE basecos > 0.9999 校验max_abs_err < 1e-3排查 eps / 权重共享
学习路径
  1. 读 4.1:理解对齐开源权重需复现 tokenizer、RoPE base、eps 与权重共享
  2. 加载开源权重并用 logits 对比,让逐层输出与 HuggingFace 对齐
  3. 对接 M8:完成 test_parity,使 max abs diff < 1e-3(cos > 0.9999)
✔ 自研实现加载权重后输出与参考实现数值一致(max abs diff < 1e-3)
核心知识点详解
  • 对齐的四要素: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 的预训练目标只有一个:「给定前面的 token,预测下一个 token」——就是交叉熵。没有别的目标,也不需要标注数据(文本本身就是标签)。所有看起来神奇的能力(推理、翻译、写代码)都是这个简单目标在超大规模下的涌现结果。理解这点,你就理解了为什么「数据质量」比「花哨的损失函数」重要得多。

从零实现一个能对齐开源权重的 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 向量)以获得更稳更大的更新。

★
对齐开源权重的实操顺序:从零实现后「对齐开源权重」不止是跑通前向,还包括:① 严格复现 tokenizer(相同 vocab 与合并规则,否则 embedding 维度对不上);② 复现 RoPE 的 base、head_dim 切分、Q/K 是否各自归一化;③ 复现 RMSNorm 的 eps 与是否含 bias;④ 用一份小语料做几步前向,检查 logits 与官方权重在固定输入下的差异 < 1e-3。能把这些对齐,你就真正读懂了这份权重的每一个字节。
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、层数、词表」三个旋钮如何共同决定模型大小形成直觉。

★
4.2 动手练习与自测:本节目标是「能自己造一个 GPT 并对齐官方权重」。
  1. 跑通:让 MiniGPT 在莎士比亚字符语料上训到 PPL < 3(约几千步、单卡即可)。判据:采样输出有英文单词与标点结构。
  2. 删行实验:删掉 Pre-LN 的 RMSNorm 或残差,观察 loss 是否发散。判据:删残差必定发散;删 RMSNorm 在深层不稳。
  3. 算参数量:V=100k、d=4096、L=32、权重共享,估算总参。参考答案:约 7B 量级(含 embedding 0.41B)。
  4. 对齐:用固定输入对比自实现与官方权重的 logits。判据:cos > 0.9999、max_abs_err < 1e-3。
  5. 采样:对比 temperature=0.2 与 1.0、top_k=1 与 50 的输出差异。判据:低温/小 top_k 更确定易重复;高温更发散。

5. Scaling Law 与规模化的现实

知识结构图 · Scaling Law 与规模化
Scaling Law 与规模化3 大知识域 · 17 个知识点
Scaling Law 演进Kaplan 幂律Chinchilla 20 tokens/param过训练 100+ tokens/param数据质量优先推理时算力 test-time scalingQwen3 约 36T tokens
学习路径
  1. 读 5.1:从 Kaplan→Chinchilla(20 tokens/param)→2026 过训练(100+)理解分配演进
  2. 用 total_cost 函数比较 70B/13B/7B 的训练与长期推理成本,理解过训练取舍
  3. 对接 M8:在 architecture.md 用参数量/FLOPs 拆解与 C≈6ND 估算串起对应关系
✔ 能讲清 Chinchilla 与过训练的差别,并算出不同规模模型的训练/推理成本
核心知识点详解
  • 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 下降明显优于低质重复数据。清出来的预算值得优先花在高质量来源上。
并行与算力估算C ≈ 6NDDP / TP / PP / SP / EPMFU 30–50%每参数约 18 字节显存ZeRO-3 每卡约 5 字节loss spike 回滚
学习路径
  1. 读 5.2:背下 C≈6ND、MFU 30–50%、每参数约 18 字节显存、ZeRO-3 每卡约 5 字节
  2. 用 C≈6ND 估算 70B/15T 算力量级并粗算 2048 卡的训练天数
  3. 对接 M8:在 architecture.md 给出模型 FLOPs 与显存估算,供手写实现核对规模
✔ 能由规模算出算力 FLOPs 与训练天数,说清 MFU 与显存分摊
核心知识点详解
  • 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.3e242048 卡约 90 天7B 按 20 tokens/param 配 140B
学习路径
  1. 完成本章自测:70B/15T 算力约 6.3e24、2048 卡约 90 天、7B 按 20 tokens/param 配 140B
✔ 自测命中判据,规模→算力→天数的换算能独立完成
核心知识点详解
  • 验证 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 2020loss 随参数量、数据量、算力呈幂律下降开启了“越大越好”的军备竞赛
Chinchilla 2022给定算力,参数与数据应同比例增长(≈20 tokens/param)纠正了“只堆参数不看数据”的浪费
2026 实践刻意过训练小模型(100+ tokens/param),能力接近大模型但推理便宜端侧与高并发场景的主流选择
数据质量优先去重、筛选、配比、合成数据的收益常大于规模数据工程成为核心竞争力
推理时算力预训练收益递减,把算力投到推理期(长思考、多次采样)催生推理模型与 test-time scaling
★
2026 年最关键的一句判断:单纯扩大预训练的边际收益在递减,推理时计算(inference-time compute)成了新的发力点:让模型在运行时多做几轮自我验证、多次采样再择优,复杂任务表现就能提升。而编排这些步骤的,正是下一阶段要学的 Harness / Context Engineering。

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 配合使用
专家并行 EPMoE 专家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} 天")
⚠
工程上比公式更重要的是:① 故障恢复:千卡集群每天都有卡挂,必须做异步 checkpoint 与弹性重启;② 数据加载不能饿死 GPU:预取、流式读取、本地缓存;③ 损失尖刺(loss spike):要有回滚到上一个 checkpoint 的机制;④ 通信重叠:计算与通信重叠做不好,MFU 会腰斩。

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 卡才放得下参数 + 优化器
⚠
估算三连:规划一次训练要算三个数:① 算力 C ≈ 6ND;② 显存 = 参数 + 梯度 + 优化器(混合精度近 18 字节/参数)÷ 分片;③ 时间 = C ÷ (卡数 × 单卡峰值 × MFU)。三个数任一算错,方案都会在集群上崩。
★
5.3 动手练习与自测:Scaling Law 的题目都是「算出来的」,不是背出来的。
  1. 算算力:70B 参数、15T token,C ≈ 6ND 是多少?参考答案:6×70e9×15e12 ≈ 6.3e24 FLOPs。
  2. 算时间:2048 张 H100(峰值 989 TFLOPS、MFU 40%)需多久?参考答案:6.3e24/(2048×989e12×0.4) ≈ 7.8e6 秒 ≈ 90 天。
  3. Chinchilla 判断:7B 模型按 20 tokens/param 应配多少数据?参考答案:约 140B token。
  4. 过训练取舍:为什么 2026 常用 100+ tokens/param?判据:推理成本主导长期成本,小模型服务更便宜。
  5. 显存估算:混合精度训练每参数显存约多少?判据:参数+梯度+Adam 状态约 18 字节/参数(未分片)。

6. FFN、归一化与残差:训练稳定性的根基

知识结构图 · FFN、归一化与残差
FFN、归一化与残差3 大知识域 · 16 个知识点
FFN 与 SwiGLUFFN 占全部参数 2/3ReLU FFN d_ff=4dGeGLU / SwiGLU 门控SwiGLU hidden=8/3 dMoE 改造 FFN
学习路径
  1. 读 6.1:理解 FFN 占总参数 2/3,SwiGLU 用 hidden=8/3·d 保持与 4d MLP 等价参数量
  2. 用代码跑参数量拆解,验算 FFN:注意力 = 8d²:4d² = 2:1
  3. 对接 M8:手写 FFN 时采用 SwiGLU,保持总参数与参考模型一致
✔ 能算清 FFN 占 2/3 参数并用 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 或维度对不上权重。
归一化与残差Post-LN → Pre-LNPre-LN 梯度直通RMSNorm 只除均方根残差缩放 ReZeroQK-Normwarmup 占总步数 1–3%
学习路径
  1. 读 6.2:理解 Post-LN→Pre-LN 的梯度直通、RMSNorm 只除均方根、残差缩放 ReZero、QK-Norm
  2. 对比 Pre-LN vs Post-LN 是否需 warmup,验证 1–3% 步数的经验值
  3. 对接 M8:手写 GPT 采用 Pre-LN + RMSNorm + 残差,配置 warmup
✔ 能说清 LN 位置/类型对训练稳定性的影响,并搭出稳定的 Pre-LN+RMSNorm 块
核心知识点详解
  • 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 先归一化防熵坍缩。两者都是「初始化稳定」路线。
稳定化前沿Muon 谱范数正交化Muon 用于 2D 权重AdamW 处理 1D 参数DeepSeek-V4 迁移 Muon
学习路径
  1. 读 6.3:理解 Muon 谱范数正交化、用于 2D 权重而 AdamW 处理 1D,DeepSeek-V4 迁移
  2. 在一个小模型上对比 AdamW 与 Muon 的收敛与吞吐,确认 Muon 的边界
  3. 完成本章自测:说明 Muon 的 weight decay 为何不能直接用 AdamW 数值
✔ 能说明 Muon 的等谱正交化机制与其对 2D 权重的适用边界
核心知识点详解
  • 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 FFNW₂·ReLU(W₁x)2·d·d_ff原版,d_ff=4d 最常见
GELU FFNW₂·GELU(W₁x)2·d·d_ffBERT / 早期 GPT
GeGLUW₂·(W₁x ⊙ GELU(W_g x))3·d·d_ff用 GELU 做门
SwiGLUW₂·(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 的隐藏维换来门控表达力,却不增加参数
ℹ
2026 落点:MoE 改造的正是 FFN:把一个大 FFN 换成 E 个小专家 + Top-K 路由,总参数量翻数倍、激活参数量却近似不变(如 DeepSeek-V4 每 token 激活约 49B)。所以「FFN 占 2/3」这件事,直接决定了 MoE 能带来多大的容量增益。

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-LNPre-LN
梯度传播需穿过每层 LN,深层易不稳残差主路无 LN,梯度直通
warmup 需求必须,且要长而精细弱得多,可省或很短
训练动态初期易 loss spike更稳定
最终性能充分训练后略优与现代配方等价,工程更稳
代表原版 TransformerLLaMA / 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 迁移的主要方向之一。

★
6.3 动手练习与自测:训练稳定性是「能不能把模型训完」的问题,比调点重要。
  1. 算比例:d=4096、ReLU FFN(d_ff=4d)、注意力 4d²,FFN 占比?参考答案:8d²/(8d²+4d²)=2/3。
  2. 解释 SwiGLU:为什么隐藏维取 8/3 d 而不是 4d?判据:门控需两个上投影(3 个矩阵),8/3 d 使参数量与 2·4d 持平。
  3. warmup 取舍:warmup 太长/太短的后果?判据:太短初期 loss spike,太长浪费算力;常取总步数 1%-3%。
  4. Pre-LN 论证:为什么 Pre-LN 对 warmup 不敏感?判据:残差主路无归一化,梯度可直通浅层。
  5. Muon 判断:它和 AdamW 分别用在哪些参数上?判据:Muon 用于 2D 权重做正交化,AdamW 用于 1D 参数。

7. KV Cache 与推理成本工程

知识结构图 · KV Cache 与推理成本
KV Cache 与推理成本3 大知识域 · 14 个知识点
KV Cache 原理单步生成降到 O(T)KV 显存公式FP8 / INT4 KV 量化memory-bound 算术强度PagedAttention 减碎片
学习路径
  1. 读 7.1:默写 KV 显存公式 KV=2·L·kv_heads·head_dim·seq·batch·dtype,理解每步降到 O(T)
  2. 用 kv_cache_gb 函数实测常见配置下的显存,理解 memory-bound 与量化/分页
  3. 对接 M8:实现 kv_cache.py 并写出增量解码,与 M2 显存公式互验
✔ 能默写并计算 KV Cache 显存,说清 memory-bound 推理与量化/分页的取舍
核心知识点详解
  • 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 已占掉大半显存。
显存测算每 token 约 0.3125 MB32K 约 10.2 GB128K 约 40.96 GBMHA 是 GQA 的 8 倍并发容量规划
学习路径
  1. 读 7.2:掌握每 token 约 0.3125 MB、32K 约 10.2 GB、128K 约 40.96 GB 的量级
  2. 手算 128K 上下文 40.96 GB 并核验 MHA 是 GQA 8 倍
  3. 对接 M8:用 KV 显存公式做并发容量规划,指导 kv_cache 实现的 batch/seq 上限
✔ 能手算 KV 显存并据此规划并发与上下文上限
核心知识点详解
  • 每 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 显存公式手算 128K ≈ 40.96 GBPagedAttention 提吞吐原因
学习路径
  1. 完成本章自测:默写 KV 公式、手算 128K,并解释 PagedAttention 为何提吞吐
✔ 自测命中判据,KV 显存公式与分页原理能独立复现
核心知识点详解
  • 默写公式与速算: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字节/元素备注
FP324几乎不用作 KV
FP162常见基线
BF162训练友好,推理常用
FP8 (E4M3/E5M2)12026 大模型推理常态
INT81量化推理
FP4 / INT40.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
✔
PagedAttention:vLLM 的 PagedAttention 借鉴操作系统虚拟内存:把 KV Cache 切成固定大小 block,按需分配、允许物理不连续,把显存碎片率从 60%-80% 降到几 %,同卡并发数可提升 2–4 倍。这是 2026 推理引擎(vLLM / SGLang)的标配。

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
★
容量规划的硬结论:KV Cache 随 序列长度 × batch 线性增长且必须常驻显存,所以「支持多长上下文」本质是显存与带宽问题,不是模型能力问题。工程上靠三件事缓解:① GQA / MQA 砍 KV 头数;② KV 量化(FP8/INT4)砍 dtype_bytes;③ 分页与卸载(PagedAttention / vLLM)把不连续显存利用起来。这也解释了为什么 2026 年几乎所有开源与闭源模型都默认 GQA。
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–8KBF16 GQA高(数十条)
长文档问答32K–128KFP8 GQA + 分页中(个位数)
超长 Agent 轨迹1MMLA + 量化 + 多卡低(需专家并行)
★
7.3 动手练习与自测:KV Cache 是面试与容量规划的高频考点,必须能手算。
  1. 默写公式:KV 显存 = ? 判据:2×L×n_kv×d_head×T×B×dtype_bytes。
  2. 手算:L=80、n_kv=8、d_head=128、T=128K、B=1、FP16。参考答案:约 40.96 GB。
  3. 对比:同配置 MHA(n_kv=64)要多少?参考答案:约 327 GB(GQA 的 8 倍)。
  4. 量化收益:FP8 量化 KV 能省多少?判据:字节减半(dtype_bytes=1),显存约减半。
  5. 工程:为什么 PagedAttention 能提吞吐?判据:减少显存碎片、提升可并发放置的序列数。

8. 长上下文工程:切分、外推与遗忘

知识结构图 · 长上下文工程
长上下文工程3 大知识域 · 14 个知识点
上下文并行CP 切序列 O(T/N)Ring Attention 环传 K/V在线 softmax 累加与 FSDP / TP 组合成 3D / 4D
学习路径
  1. 读 8.1:理解上下文并行切序列 O(T/N),Ring Attention 用在线 softmax 跨块累加
  2. 梳理 CP/SP/TP/EP 各自切分轴,说明与 FSDP/TP 的 3D/4D 组合
  3. 完成本章自测:解释 CP 解决什么、1M 单卡 KV 约 320 GB
✔ 能说清 CP/Ring Attention 去显存的机制及其与 FSDP/TP 的组合
核心知识点详解
  • 上下文并行切序列 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 缩短有效长度关键证据重排首尾有效上下文长度Needle-in-a-Haystack 局限
学习路径
  1. 读 8.2:理解 Lost in the Middle 现象与注意力被长序列稀释的机制
  2. 跑一次 RAG 或关键证据重排实验,观察中段遗忘与首尾优先
  3. 完成本章自测:列出三种缓解中段遗忘的办法
✔ 能复现中间遗忘并用 RAG/重排缓解,识别有效上下文长度与 NiH 局限
核心知识点详解
  • Lost in the Middle 现象:长上下文中模型对中间位置的证据召回与利用明显弱于首尾。越长的序列稀释越严重,是 RAG 与重排要解决的问题。
  • 三种缓解办法:① RAG 只在上下文里放最相关文本,缩短有效长度;② 关键证据重排到开头/结尾(黄金位);③ 用注意力/位置增强(如延长训练、中间采样)。
  • 有效上下文长度 ≠ 标称长度:模型可能「读得进」长文本(不丢单条信息)但「用不好」(跨位置推理弱)。要分别测召回与多跳利用,别只看能塞进去的长度。
  • Needle-in-a-Haystack 的局限:只测「单条关键信息能否召回」,不测长程推理与累积,通过 NiH 不代表上下文能力强。常见坑:用 NiH 全通过宣称长上下文好,业务长程任务照样崩。
动手自测要点解释 CP 解决什么1M 单卡 KV 约 320 GB三种缓解中段遗忘办法
学习路径
  1. 完成本章自测:解释 CP 解决什么,手算 1M 单卡 KV 约 320 GB,梳理三种中段遗忘缓解
✔ 自测命中判据,能同时讲透长上下文并行与中段遗忘
核心知识点详解
  • 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 上下文的可行性
ℹ
2026 落地:Ring Attention 是 1M token 级训练/推理的底层支柱之一。工程上常与 ZeRO / FSDP(切参数)、TP(切单层) 组合成 3D/4D 并行;DeepSeek、Qwen 等的超长上下文训练配方里,CP + Ring 几乎是默认配置。

8.2 长上下文的「中间遗忘」现象

一个反直觉的事实:把关键信息放在超长 prompt 的中段,模型表现往往比放在开头或结尾更差——这被称为「Lost in the Middle」(Liu 等,2023–2024)。两条主因:① 注意力在长序列里被稀释,模型对中段 token 分配的权重偏低;② 预训练与 SFT 的监督信号大多集中在文档首尾(如段落首尾、对话头尾),模型「习惯」关注两端。

缓解手段原理代价
检索增强(RAG)只把相关片段送入上下文,缩短有效长度需检索质量高,否则引入噪声
重排序 / 重打包把关键证据移到首尾需先识别关键片段
长上下文续训用长样本继续预训练,弥合位置外推算力最贵但最可靠
压缩 / 摘要中段先摘要再拼回可能丢细节
位置插值 + 温度缩放缓解 RoPE 外推衰减过长仍有损
⚠
评测要诚实:很多「支持 128K / 1M」的宣称来自「针在草堆里」(needle-in-a-haystack)式测试,但真实长文档问答远难于此。中段信息检索、跨段推理、长程依赖一致性才是真考验。上线前务必用你自己领域里的长文档任务实测,而不是只看榜单的针测试分数。
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多任务长上下文仍需按领域实测
有效上下文长度能答对的最长长度标称窗口常远大于有效窗口
跨段一致性多段信息整合推理中段证据易被忽略
⚠
上线前必做:用你自己领域的长文档任务实测:中段信息检索、跨段推理、长程一致性。把「针测试满分」当成能力证明,是 2026 依然常见、也依然致命的误判。
★
8.3 动手练习与自测:长上下文题考「切分 + 外推 + 评测」三件事。
  1. 解释 CP:上下文并行解决什么问题?判据:把序列摊到多卡,显存从 O(T) 降到 O(T/N)。
  2. Ring 机制:为什么 Ring Attention 能保持注意力完整?判据:K/V 块沿环传递 + 在线 softmax 累加。
  3. 算收益:1M 上下文、80 层、GQA 8 头、head_dim 128,单卡 KV 约多少?参考答案:约 320 GB,故必须多卡 CP / 分页 / 量化。
  4. 诊断:模型「支持 128K」但长文档中段问题答错,最可能原因?判据:Lost in the Middle + 外推衰减,非单纯显存问题。
  5. 对策:给出三种缓解中段遗忘的办法。判据:RAG 缩长、关键证据重排到首尾、长上下文续训 + 位置缩放与温度。

项目里程碑

贯穿项目 · Hamauls Orion
M8 从零实现 GPT 并对齐开源权重 第 39–46 周

Hamauls Orion 的模型层自己实现:手写 Multi-Head Attention、RoPE、RMSNorm、SwiGLU、KV Cache,并做到能加载开源权重、输出与 HuggingFace 参考实现数值一致。这一步打通「架构 ↔ 权重 ↔ 推理」的完整链条。

本阶段产出(直接进入项目仓库)
验收标准:自研实现的输出与参考实现数值一致;能白板画出数据流并算清参数量、FLOPs 与 KV Cache 显存;对「为什么用 RoPE 而不是绝对位置编码」能给定量解释。

阶段练习项目

PROJECT 1
MiniGPT 从零训练
在字符级或小词表语料上从零训练一个 GPT,让它能生成通顺片段。要求:手写全部组件(不用 nn.TransformerEncoder),记录 PPL 曲线,并做 Top-K / 温度采样对比。
要达成的效果
  • 在字符级 / 小词表语料上从零训到 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
边界 · 不做

不做真实大规模语料与多卡并行训练,仅做字符级 / 小词表的能力验证。

PROJECT 2
KV Cache 与 GQA 实测
实现 MHA / GQA / MQA 三种注意力,测量在 4K / 32K 上下文下的显存占用与生成吞吐,画出对比表并给出部署建议。
要达成的效果
  • 在 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 基线下的相对差异。

PROJECT 3
长上下文外推实验
在 4K 训练、8K 推理的条件下,对比「直接外推」「位置插值」「NTK 缩放」「YaRN」四种方案在困惑度上的表现差异。
要达成的效果
  • 在固定 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(含有效窗口测量与选型建议)
边界 · 不做

不做长上下文续训 / 退火微调优化外推,仅评估固定基座下外推方法的相对表现。

PROJECT 4
技术报告精读笔记
选 DeepSeek-V4 或 Qwen3.5 的技术报告,按「架构 / 数据 / 后训练 / 基建 / 评测」五栏做一份结构化笔记,并标出与 2017 原版的每一处差异及其理由。
要达成的效果
  • 产出「架构 / 数据 / 后训练 / 基建 / 评测」五栏结构清晰的精读笔记
  • 逐条标出与 2017 原版 Transformer 的差异并给出动机,条数不少于 8 条
  • 笔记能直接支撑选型与后续手写实现决策(含关键数值与引用)
功能需求
  • 选定 DeepSeek-V4 或 Qwen3.5 并通读其技术报告
  • 五栏各提取关键数字(参数量 / 层数 / 词表 / RoPE base / 数据量 / 超参)并标注来源引用
  • 单列一栏逐条比对与 2017 原版的差异(如 Pre-LN/RMSNorm/GQA/RoPE/SwiGLU/MLA/MoE)并给理由
  • 在「评测」栏给出作者自评与第三方独立评估的差距对照
交付物
  • report_notes.md 结构化精读笔记
边界 · 不做

不做复现训练,只做文献精读与要点提取。

PROJECT 5
从零实现 GPT 并对齐开源权重(M8)
不依赖 nn.TransformerEncoder,从 embedding / RoPE / RMSNorm / GQA / SwiGLU / MoE 全部手写一个 Decoder-only GPT;先在小语料上训到会续写,再严格复现某开源权重(tokenizer、RoPE base、eps、权重共享)使固定输入下的 logits 差异 < 1e-3。这是本阶段里程碑 M8 的交付物。
要达成的效果
  • 无需 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 / 对齐要点)
边界 · 不做

不做多轮 / 指令能力级预训练,不对齐闭源权重(只选公开权重),不做推理部署与量化。

常见误区

面试高频问题速答

写出注意力的公式并解释每个符号。

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),改动小、成本低,但过长时质量衰减;② 稀疏/滑动窗口注意力降低计算量,但可能丢远距离依赖;③ 直接用长样本继续预训练(最可靠也最贵)。实践中常组合使用:先做位置插值让模型「能跑」,再用少量长样本文本做续训让模型「跑得好」,并配合注意力温度缩放。

学习资源

Attention Is All You Need(原论文)论文 arxiv.org/abs/1706.03762 起点。2026 年重读一遍,你会发现现代模型改了哪些、又保留了哪些。 The Annotated Transformer教程 nlp.seas.harvard.edu/annotated-transformer/ 逐行注释实现,理解 Transformer 的最佳入门材料。 nanoGPT(Karpathy)GitHub github.com/karpathy/nanoGPT 极简可训练 GPT 实现,本阶段项目的直接参考。 FlashAttention 论文与实现论文 github.com/Dao-AILab/flash-attention 理解 IO 感知与在线 softmax,读第 2 节即可掌握核心思想。 RoFormer(RoPE 原论文)论文 arxiv.org/abs/2104.09864 旋转位置编码的源头,配合 YaRN 一起读。 DeepSeek-V3 技术报告论文 arxiv.org/abs/2412.19437 MLA + 细粒度 MoE + 无辅助损失负载均衡,是理解现代架构性价比的必读。 Chinchilla(Training Compute-Optimal LLMs)论文 arxiv.org/abs/2203.15556 算力最优分配的经典结论。 How Frontier Labs Train LLMs(2026 综述)长文 jxzhangjhu.github.io/blog/2026/how-frontier-labs-train-llms 横向对比各家的架构与训练配方,看清共识与分歧。 Lilian Weng · The Transformer Family长文 lilianweng.github.io/posts/2023-01-27-the-transformer-family-v2/ 把 Transformer 家族的变体梳理得非常清楚。
★
2026 形势提示:2026 年的架构共识已经成型:现代 Decoder 块 + 细粒度 MoE + GQA + RoPE + RMSNorm/SwiGLU + FlashAttention + Muon。真正还在激烈竞争的是三件事:1M+ 长上下文(混合 / 线性注意力)、推理期算力(推理模型)、Agentic RL(长程训练稳定性)。把本阶段的三个「为什么」讲透,后面两个阶段你会跑得很快。