TransformerBlock 是 Decoder-only 大语言模型的核心计算单元。Embedding 负责把 token ID 映射成向量,LM Head 负责把最终隐藏状态映射回词表,而模型绝大部分参数、计算量和 KV Cache 都位于一层层重复堆叠的 TransformerBlock 中。
TransformerBlock 在模型中的位置
一个完整的 Causal Language Model 可以拆成
graph LR
Tokens("Token IDs [B, S]") --> Embedding("Token Embedding [B, S, H]")
Embedding --> Block0("TransformerBlock 0")
Block0 --> Block1("TransformerBlock 1")
Block1 --> More("...")
More --> BlockN("TransformerBlock L-1")
BlockN --> Norm("Final RMSNorm")
Norm --> Head("LM Head")
Head --> Logits("Logits [B, S, V]")
这里
| 配置 | 符号 | 数值 |
|---|---|---|
| Hidden size | 896 | |
| TransformerBlock 数量 | 24 | |
| Query head 数量 | 14 | |
| Key/Value head 数量 | 2 | |
| Head dimension | 64 | |
| MLP intermediate size | 4864 | |
| Vocabulary size | 151936 |
Embedding 输出的 shape 是
Block 的总体结构
Qwen 使用 Pre-Norm 结构。一个 Block 可以写成两个连续的残差子层。
第一条路径负责 token 之间的信息交换,第二条路径负责每个 token 内部的非线性特征变换。Attention 和 MLP 都不直接替换输入,而是通过残差连接把增量写回主干。
flowchart LR
X["x"] --> Norm1["RMSNorm"]
Norm1 --> Attention["Causal Self-Attention"]
X --> Add1["+"]
Attention --> Add1
Add1 --> U["u"]
U --> Norm2["RMSNorm"]
Norm2 --> MLP["SwiGLU MLP"]
U --> Add2["+"]
MLP --> Add2
Add2 --> Y["y"]
对应代码位于 layers.py。第一次残差保存
Block 原始输入,第二次残差保存 Attention 输出与原输入相加后的结果。
1 | residual = hidden_states |
两次残差不能混用。如果 MLP 输出错误地加回 Block 最初的输入,网络结构就不再等价于 Qwen 的 Decoder Layer,即使所有 Linear 权重完全一致,最终 logits 也会产生偏差。
RMSNorm
RMSNorm 根据一个 token 隐藏向量的均方根进行归一化。设输入
RMSNorm 与 LayerNorm 的主要区别是它不减去均值,也不使用可训练 bias。LayerNorm 同时控制均值与方差,RMSNorm 只控制向量整体尺度,因此计算路径更短。
项目中的 RMSNorm
先把输入转换成 float32,再进行平方、Reduction 和平方根倒数运算。
1 | input_dtype = hidden_states.dtype |
模型权重可能使用 BF16,但 Reduction 会累加
RMSNorm 包含逐元素平方、Reduction、rsqrt
和逐元素缩放。它通常是 memory-bound
算子,优化重点不是增加复杂计算,而是减少全局内存访问、让读取连续合并,并在一个
Kernel 内完成 Reduction 与归一化,避免保存不必要的中间 Tensor。
Rotary Position Embedding
Self-Attention 本身只计算 token 向量之间的相似度。如果调换两个 token 的位置而不提供位置信息,Attention 无法区分它们的先后顺序。RoPE 不把位置向量直接加到 hidden states,而是根据位置旋转 Query 和 Key,使它们的点积自然携带相对位置信息。
频率构造
设每个 head 的维度为
位置
Qwen2.5 使用
RotaryEmbedding
使用 float32 生成角度、cosine 和 sine,最后转换回 hidden states 的
dtype。位置越长,低精度浮点数越难准确表示角度,因此频率计算不能轻易降到
BF16。
旋转操作
项目采用 Qwen 的 split-half 旋转形式。把最后一个维度拆成两半:
定义:
最终旋转为:
RoPE 只作用于 Query 和 Key,不作用于 Value。位置应影响“当前 Query 应该关注哪个 Key”,而 Value 负责承载被读取的内容。
Query 的 shape 是
Grouped-query Attention
Q/K/V Projection
Attention 输入是
对于 Qwen2.5-0.5B,
以
| Tensor | Projection 后 | 调整 head 维度后 |
|---|---|---|
| Query | ||
| Key | ||
| Value |
项目实现中 Q/K/V projection 使用 bias,O projection 不使用 bias,MLP
的三个 projection 也不使用 bias。这些细节必须与目标模型完全一致,否则
state_dict 即使能加载,计算也无法对齐。
GQA 的共享方式
传统 Multi-head Attention 为每个 Query head 配置独立的 Key 和 Value
head,因此通常有
每个 K/V head 服务的 Query head 数量为:
repeat_key_value
将 expand 形成
expand
表达的是共享关系,不会像直接复制那样立刻分配七倍底层存储。真正的 KV
Cache 仍只需要保存 2 个 K/V heads,这正是 GQA 对 Decode 性能的价值。
KV Cache 成本
每一层、每个 token 的 KV Cache 元素数量为:
前面的 2 分别代表 Key 和 Value。Qwen2.5-0.5B 使用 BF16,每个元素占 2 bytes,因此每层每 token 的理论存储量为:
24 层合计约为:
如果使用 14 个 K/V heads 的传统 MHA,同样计算约为 84 KiB/token,正好是 GQA 的 7 倍。这里没有计入 allocator 对齐、元数据和框架对象开销,但可以展示 GQA 为什么适合长上下文推理。
Causal Attention
因果约束
Decoder-only 模型生成当前位置 token 时不能读取未来 token。长度为 4 时,允许访问的位置如下。
| Query 位置 | 可读取的 Key 位置 |
|---|---|
| 0 | 0 |
| 1 | 0, 1 |
| 2 | 0, 1, 2 |
| 3 | 0, 1, 2, 3 |
build_causal_attention_mask
创建 shape 为
Scaled Dot-product Attention
应用 RoPE 并扩展 K/V heads 后,Attention score 为:
这里除以
随后计算概率和输出:
score 和概率的 shape 是
softmax 强制使用 float32:
1 | attention_weights = torch.softmax( |
指数运算对数值范围敏感,BF16/FP16 更容易发生精度损失。未来替换为融合 Attention Kernel 时,必须把这种累积和归一化精度纳入 correctness 标准。
SwiGLU MLP
Attention 负责 token 之间的信息交换,MLP 负责对每个 token 的特征独立进行非线性变换。Qwen 使用 SwiGLU,而不是简单的两层 Linear 加 ReLU。
其中:
输入和输出 shape 都是
SwiGLUMLP
的实现保持三个 Linear 层彼此独立,因为它们对应 Hugging Face 权重中的
gate_proj、up_proj 和
down_proj。生产 Kernel 可以把 Gate 和 Up projection
合并成一次更宽的 GEMM,但 reference
实现首先追求结构透明和权重一一对应。
参数分布
Qwen2.5-0.5B 的单个 TransformerBlock 大约包含 1491 万参数。Attention、MLP 和 Norm 的参数量可以按矩阵 shape 直接计算。
| 组件 | 参数量 |
|---|---|
| Q/K/V/O projection 与 Q/K/V bias | 1,836,160 |
| Gate/Up/Down projection | 13,074,432 |
| 两个 RMSNorm | 1,792 |
| 单个 Block 合计 | 14,912,384 |
24 个 Block 约包含 3.58 亿参数。Embedding 与绑定的 LM Head 约包含 1.36 亿参数,因此总量接近 4.94 亿,符合 0.5B 模型的命名。参数分布也说明 MLP GEMM 是模型 FLOPs 的重要来源,而 Decode 阶段的 Attention 还要持续读取不断增长的 KV Cache。优化 TransformerBlock 不能只关注单一算子,需要区分 Prefill 和 Decode 的瓶颈。
Prefill 与 Decode 中的 Block
Prefill
Prefill 一次输入完整 prompt。若 prompt 长度为
Attention score 的 shape 是
Decode
Decode 每轮只输入一个新 token,因此当前输入 shape 是
每个生成 token 都必须依次经过 24 个 Block。Decode 中的小 GEMM、Kernel launch 和 KV Cache 读取更突出,通常比 Prefill 更容易受到内存带宽和调度开销限制。
每层独立缓存
KV Cache 不是整个模型共享的一份缓存。每个 Block 的投影权重不同,产生的 K/V 表示也不同,因此模型需要维护 Layer 0 到 Layer 23 各自独立的 KV Cache。
现有 engine/generation.py
把 Hugging Face 返回的 past_key_values
作为整体状态保存,内部实际上包含每一层的 K/V。Prefill 建立缓存,Decode
每轮读取并更新它。
总结
TransformerBlock 是原先 Hugging Face Qwen 模型内部
model.model.layers[i] 的显式实现。每个 Block 通过
RMSNorm、RoPE、Grouped-query Attention、SwiGLU MLP
和两次残差连接,把输入
理解这一层的 shape、数值精度、KV Cache 和 Prefill/Decode 行为,是继续学习 LLM Runtime 与 GPU Kernel 的连接点。上层 Runtime 决定请求和缓存如何组织,底层 Kernel 决定 Block 中每一次 Reduction、GEMM、Softmax 和内存访问如何高效执行。
参考资料
Hugging Face | Qwen/Qwen2.5-0.5B-Instruct
GitHub | Euler0525/ai-infra-learning
附录
TransformerBlock 对应关系
| Hugging Face Qwen2 | 项目参考实现 |
|---|---|
Qwen2RMSNorm |
RMSNorm |
Qwen2RotaryEmbedding |
RotaryEmbedding |
Qwen2Attention |
CausalSelfAttention |
Qwen2MLP |
SwiGLUMLP |
Qwen2DecoderLayer |
TransformerBlock |
Qwen2Model |
QwenDecoder |
Qwen2ForCausalLM |
TorchReferenceQwen |
讨论
评论