CS336(三):现代 LLM 的基础架构
今天的大语言模型仍然建立在 Transformer 上,但真正落到模型结构里,Norm、激活函数、位置编码和 Attention 都已经和最初版本有了不少差别。
这些改动主要围绕三个问题:怎样让训练更稳定,怎样控制参数量,以及怎样减少推理时的计算和显存开销。
1. 归一化与残差路径
1.1 Pre-Norm
Pre-Norm 先对输入做归一化,再送入 Attention 或 FFN,最后加上原始输入。用 表示 Attention 或 FFN,则 ;Post-Norm 则是 。
flowchart LR
X[x] --> N[Norm]
N --> F[Attention / FFN]
F --> A[Add]
X --> A
A --> Y[output]区别在于残差连接是否经过 Norm。Pre-Norm 保留了从输入直接到输出的路径,信息和梯度都可以沿这条路径传递。Post-Norm 则在相加后再做归一化,梯度仍要经过 Norm。模型较深时,Pre-Norm 通常更容易稳定训练,也更少出现梯度峰值。
从梯度看,Pre-Norm 的导数中有一个直接来自 的恒等项 ;Post-Norm 则连这条路径也要乘上 Norm 的导数。这里要保留的是残差支路的恒等连接,与 FFN 使用 ReLU 还是其他激活函数无关。
1.2 LayerNorm 与 RMSNorm
对一个有 个元素的隐藏向量,LayerNorm 先减均值,再除以标准差;RMSNorm 保留尺度归一化,省掉去均值和可学习偏移项:
其中 和 沿隐藏维度计算, 是可学习的缩放参数。RMSNorm 简化了计算,在不少语言模型中可以取得与 LayerNorm 接近的效果。
Norm 和 Softmax 的计算量不大,却要反复读取数据、计算统计量,再写出结果。它们往往慢在访存上。下面这组实验就能看出这种差别:
算子类型 | FLOP 占比 | 运行时间占比 |
|---|---|---|
张量收缩(如矩阵乘法) | 99.80% | 61.0% |
归一化 | 0.17% | 25.5% |
逐元素操作 | 0.03% | 13.5% |
归一化只占很少的 FLOP,却花掉了四分之一左右的时间。优化这类算子,重点是减少数据读写。
很多 LLM 也去掉了 FFN 等线性层的 Bias。这能简化结构,在一些模型实验中没有明显损失效果,还能改善训练稳定性。
2. FFN 与激活函数
2.1 GELU
GELU 定义为 ,其中 是标准正态分布的累积分布函数。ReLU 会把负数直接置零,GELU 则让负输入平滑地衰减:接近零时仍有输出,向负无穷移动时逐渐趋近于零。整个函数连续可导。

2.2 门控:GLU、GeGLU 与 SwiGLU
门控 FFN 把输入分别乘以两组权重,再将结果逐元素相乘:。其中 表示逐元素乘, 是激活函数。一条分支的输出会放大或减弱另一条分支的特征。
例如,普通 ReLU 分支输出为 ,加入另一条分支 后,结果变成 。若某个特征的 ,那么 时输出为 , 时输出为 。这让特征的强弱能随输入调整; 也可以为负,因此它不是一个固定在 0 到 1 之间的开关。
变体 | 门控激活 | 中间表示 |
|---|---|---|
GLU | Sigmoid | |
ReGLU | ReLU | |
GeGLU | GELU | |
SwiGLU | Swish |
这里采用 。相乘之后还需要输出投影 ,把中间维度映射回模型维度。
普通 FFN 有输入、输出两个矩阵,参数量约为 ;门控 FFN 有三个矩阵,约为 。若普通 FFN 使用 ,要维持接近的参数量,就令:
门控 FFN 的中间宽度因此取普通 FFN 的 ,约为 。实现时通常会取整,方便 GPU 计算。
2.3 Attention 与 MLP 并行
通常一个 Block 先算 Attention,再把结果送入 MLP。并行结构让两者处理同一份输入,最后把两路结果加回残差:。两条分支没有前后依赖,可以共享 Norm,并把部分矩阵乘法合并执行;实际能加速多少取决于实现。
3. 位置编码与注意力结构
3.1 RoPE
RoPE 将 Q、K 各自的 维向量拆成 组二维向量,每组分别旋转。第 个位置上的第 组旋转 ,其中 从 0 开始,常用 , 是频率基数。最初的 RoPE 取 。
例如 时,共有 4 组二维向量,对应的 分别为 、、、。位置每向后移动一个 token,各组就分别多转这么多弧度。越靠后的维度组,频率越低,转得越慢。
同一个频率组内,两个位置旋转后的内积满足:
Q、K 各自按所在位置旋转,内积中的旋转角度最终取决于两个位置的差 。例如,we know 出现在位置 0、1,或出现在位置 2、3,位置差都是 1,因此这一对位置引入的相对旋转相同。

RoPE 沿用了正弦位置编码按维度设置不同频率的思路。正弦位置编码把 、 组成的向量加到词向量上;RoPE 则用它们构造旋转矩阵,直接旋转 Q/K,使内积带上相对位置。
不同频率的作用可以从旋转速度理解。高频组对一个 token 的位置变化就很敏感,适合区分“当前位置”“前一个位置”等关系;低频组在一段距离内转动很小,Q/K 的内容匹配较少被位置变化打乱,更适合保留语义关系。高频用于形成位置注意力,低频作为语义信息通道。
3.2 Attention 的计算量与 KV Cache
Attention 的开销可以分成两部分:做多少次计算,以及从显存读写多少数据。
Prefill
从输入到 Q、K、V
设一批有 条序列,每条包含 个 token,模型隐藏维度为 ,输入 的 Shape 为 。注意力有 个头,按每头维度为 的标准 MHA 计算。
输入分别乘三个投影矩阵,得到 、、。
以 Q 为例,计算量为 。K、V 各做一次同样的投影。多头注意力的结果拼接后,还要乘输出投影矩阵 。按一次乘法和一次加法计 2 FLOP,这四次投影合计 FLOP,忽略常数后为 。
token 之间的注意力计算
设每条序列有 个 token,每个 token 的 个特征分给 个头,因此每个头的 Q、K、V 形状均为 。
一个头先算 ,矩阵形状为 。乘上 个头和 条序列,计算量为:
分数经过 Softmax 得到注意力权重后,还要对 V 加权求和。每个头的矩阵形状为 ,也是同样的计算量。因此, 和对 V 加权求和合计约 FLOP,量级为 。固定总维度 时,头越多,每个头处理的特征越少,因此头数 在计算量中约掉了。两类计算合起来,Attention 的主要计算量为 :前一项来自投影,后一项来自 token 之间的注意力计算。
数据量
朴素实现会显式保存每个头的分数矩阵,主要涉及三类数据:
数据 | 元素数量的计算过程 | 合计 |
|---|---|---|
Q、K、V | 每个张量的 Shape 为 ,包含 个元素;Q、K、V 共 3 份 | |
各头的 注意力分数 | 每条序列的每个头产生一个 分数矩阵;一共 条序列、每条 个头,因此有 个元素 | |
四个投影矩阵 | 各为 ,每个包含 个参数;共 4 个矩阵,由所有序列和 token 共享 |
忽略常数后,数据规模合计为 。
Decode:KV Cache
假设已经处理完前 100 个 token,现在处理第 101 个。掩码注意力使旧 token 看不到后来新增的 token,因此前 100 个 token 的 K/V 不会随之改变。把这些结果保存下来,就是 KV Cache。
访存量:设 Batch Size 为 ,隐藏维度为 ,当前已有 个历史 token。按标准 MHA 计算,每个 token 的 K、V 各有 个元素。
数据 | 计算过程 | 元素数量 |
|---|---|---|
读取投影权重 | 各为 ,一批请求共用这 4 个矩阵 | |
读取历史 K/V | 每条请求有 个历史 token,每个 token 保存 K、V 两份 维向量,共 条请求 | |
写入新 K/V | 每条请求只计算新 token 的 Q/K/V,将其中 K、V 两份 维向量追加到缓存 |
从前缀开始连续生成 步,投影权重每步读取 个元素,累计为 ;历史 K/V 的读取量随长度增长,累计为 ;新 K/V 每步写入 个元素,累计为 。三项相加:
忽略常数及低阶项,生成 步的主要累计访存量为 。
计算量: 步投影累计为 ;注意力计算随历史长度累加,为 。因此累计计算量为 。
整段输入计算时,同一份权重可供 个 token 使用;逐 token 生成时,每步只有 个新 token 使用权重,但每个新 Query 都要再读一遍历史 K/V。因此生成时的数据复用更少,计算密度下降。MQA、GQA 通过进一步减少 KV 头数,降低这部分缓存和读取量。
3.3 MHA、MQA、GQA 与 MLA
结构 | K/V 的组织方式 | 对缓存的影响 |
|---|---|---|
MHA | 每个 Q 头有对应的 K/V 头 | 保存完整的多头 KV |
MQA | 所有 Q 头共享一个 K/V 头 | 大幅减少 KV 头数 |
GQA | 每组 Q 头共享一个 K/V 头 | 在 MHA 与 MQA 之间折中 |
MLA | 将 KV 信息压到低维潜变量 | 缓存低维表示,结合投影完成注意力计算 |
MLA 缓存的是压缩后的表示。计算注意力时,可以将部分恢复投影合并到其他矩阵运算中,避免先展开全部历史 KV,再把它们重新存回显存。
以标准 MHA 为例,一层为 个请求、每个请求 个历史 token 保存 K/V,需要 个元素。MQA 把 KV 头数降为 1,变成 ;GQA 若使用 个 KV 头,则是 。Q 仍然有 个头,减少的是历史 KV 的存储和读取。
3.4 滑动窗口与全局注意力
滑动窗口注意力只读取当前 token 附近的一段历史,序列再长,也只计算窗口内的注意力。
例如,可以把三层带 RoPE 的滑动窗口注意力和一层 NoPE 全局注意力放在一组。NoPE 表示不额外加入显式位置编码。前三层只看附近的 token,减少计算;第四层再看完整的已知上下文,让距离较远的 token 也能交换信息。
4. 超参数与训练稳定性
4.1 宽度、深度和词表
普通 FFN 的中间维度通常取 。门控 FFN 多一组投影,为保持参数量接近,中间宽度缩为普通 FFN 的 ,约为 。
注意力头数 与每头维度 通常满足 ,即各头的总维度等于模型隐藏维度。部分模型会使用更大的 Attention 内部维度。
模型的宽深比为 ,其中 是隐藏维度, 是层数。一般来说,模型宽深比为 128 时表现最好。此外,多卡并行计算时,张量并行拆维度,流水线并行拆层数,宽度和层数的选择还取决于卡间通信速度。
词表大小通常为单语言 30K~50K、多语言 100K~250K。
4.2 权重衰减与 Dropout
大规模预训练中,权重衰减的主要作用是改善优化过程、降低训练 Loss,而非防止过拟合。
由于预训练通常只遍历一个 Epoch,较少重复使用同一批样本,因此很多模型减小或去掉 Dropout。
4.3 z-loss 与 QK Norm
训练后期,部分模型会出现输出 Logit 数值过大的问题。z-loss 通过加入 惩罚项,以约束 Logit 的数值尺度,提高训练稳定性。当各项 Logit 之间的差距增大时,Softmax 概率会更集中。
Q/K 范数持续增大会放大注意力分数,使 Softmax 越来越极端。QK Norm 在计算分数前对 Q/K 做归一化,可以控制它们的尺度,稳定注意力计算。

