Transformer 架构:从 Attention 到完整前向传播¶
更新日期:2026-04-25
TL;DR
一个 Decoder-Only Transformer(LLaMA 风)= Embed → N × (Pre-Norm → Attn + Residual → Pre-Norm → SwiGLU + Residual) → Norm → LM Head。
参数量 ≈ 12 L D²(忽略 embedding),训练 FLOPs ≈ 6 × N × tokens。
长序列的 O(S²) 瓶颈来自 Attention 的 QK 矩阵乘;短序列则是投影 + FFN 主导。
一、整体架构¶
flowchart LR
tok["token_ids<br/>[B, S]"] --> emb["Embedding<br/>[B, S, D]"]
emb --> L1["Block × N<br/>(Pre-Norm + Attn + FFN)"]
L1 --> norm["RMSNorm<br/>[B, S, D]"]
norm --> head["LM Head<br/>(tied with Embed)"]
head --> logits["logits<br/>[B, S, V]"]
classDef data fill:#f5f3eb,stroke:#bdb9ab,color:#1a1a1a;
classDef op fill:#fff,stroke:#cc785c,color:#1a1a1a;
class tok,emb,norm,logits data
class L1,head op
主要形状符号:B = batch,S = seq_len,D = d_model,H = n_heads,d_h = D/H,F = d_ff,V = vocab_size,L = n_layers。
from jaxtyping import Float, Int
from torch import Tensor, nn
class TransformerLM(nn.Module):
def __init__(self, V: int, D: int, L: int, H: int, F: int):
super().__init__()
self.embed = nn.Embedding(V, D)
self.layers = nn.ModuleList([TransformerBlock(D, H, F) for _ in range(L)])
self.norm = RMSNorm(D)
self.lm_head = nn.Linear(D, V, bias=False)
# weight tying — LLaMA & GPT 惯例,省 V×D 参数
self.lm_head.weight = self.embed.weight
def forward(self, ids: Int[Tensor, "B S"]) -> Float[Tensor, "B S V"]:
x: Float[Tensor, "B S D"] = self.embed(ids)
for layer in self.layers:
x = layer(x)
return self.lm_head(self.norm(x))
二、单个 Transformer Block¶
Pre-Norm 是现代主流(LLaMA / GPT-2+ / Mistral),Post-Norm 是原始 "Attention Is All You Need" 的做法。
flowchart LR
x["x<br/>[B,S,D]"] --> n1[RMSNorm]
n1 --> attn[Causal<br/>Self-Attn]
attn --> r1((+))
x -.residual.-> r1
r1 --> n2[RMSNorm]
n2 --> ffn[SwiGLU<br/>FFN]
ffn --> r2((+))
r1 -.residual.-> r2
r2 --> y["x′<br/>[B,S,D]"]
classDef data fill:#f5f3eb,stroke:#bdb9ab,color:#1a1a1a;
classDef op fill:#fff,stroke:#cc785c,color:#1a1a1a;
class x,y data
class n1,n2,attn,ffn op
class TransformerBlock(nn.Module):
def __init__(self, D: int, H: int, F: int):
super().__init__()
self.attn_norm = RMSNorm(D)
self.attn = CausalSelfAttention(D, H)
self.ffn_norm = RMSNorm(D)
self.ffn = SwiGLU_FFN(D, F)
def forward(self, x: Float[Tensor, "B S D"]) -> Float[Tensor, "B S D"]:
x = x + self.attn(self.attn_norm(x))
x = x + self.ffn(self.ffn_norm(x))
return x
Pre-Norm vs Post-Norm¶
| 维度 | Pre-Norm(现代) | Post-Norm(原始) |
|---|---|---|
| 结构 | x + f(norm(x)) |
norm(x + f(x)) |
| 残差路径 | 恒等(无 norm),梯度直通 | 残差被 norm 缩放 |
| 训练稳定性 | 好,可不加 warmup 堆到 100+ 层 | 需要 lr warmup,深层易 collapse |
| 最终输出 | 需要一层最终 RMSNorm | 不需要 |
| 代表模型 | LLaMA、Mistral、GPT-2+、Qwen | 原始 Transformer、BERT |
Pre-Norm 胜出的核心原因:残差路径恒等让深层梯度不衰减。参考 On Layer Normalization in the Transformer Architecture (2020)。
三、Self-Attention 完整推导¶
3.1 八步流程¶
flowchart LR
x["x [B,S,D]"] --> proj["1. W_q / W_k / W_v<br/>(可融合成 W_qkv)"]
proj --> split["2. 拆多头<br/>[B,H,S,d_h]"]
split --> rope["3. RoPE 旋转<br/>(只作用 Q,K)"]
rope --> score["4. Q K^T / √d_h<br/>[B,H,S,S]"]
score --> mask["5. Causal Mask<br/>(上三角 → -inf)"]
mask --> sm["6. softmax"]
sm --> av["7. attn · V<br/>[B,H,S,d_h]"]
av --> merge["8. reshape + W_o<br/>[B,S,D]"]
classDef data fill:#f5f3eb,stroke:#bdb9ab,color:#1a1a1a;
classDef op fill:#fff,stroke:#cc785c,color:#1a1a1a;
class x data
class proj,split,rope,score,mask,sm,av,merge op
from torch.nn.functional import scaled_dot_product_attention as sdpa
class CausalSelfAttention(nn.Module):
def __init__(self, D: int, H: int):
super().__init__()
assert D % H == 0
self.H, self.d_h = H, D // H
# 融合 QKV 投影,省一次 launch + 更好的显存布局
self.W_qkv = nn.Linear(D, 3 * D, bias=False)
self.W_o = nn.Linear(D, D, bias=False)
def forward(
self,
x: Float[Tensor, "B S D"],
positions: Int[Tensor, "S"],
) -> Float[Tensor, "B S D"]:
B, S, D = x.shape
qkv: Float[Tensor, "B S 3D"] = self.W_qkv(x)
q, k, v = qkv.chunk(3, dim=-1)
# [B,S,D] → [B,H,S,d_h]
q = q.view(B, S, self.H, self.d_h).transpose(1, 2)
k = k.view(B, S, self.H, self.d_h).transpose(1, 2)
v = v.view(B, S, self.H, self.d_h).transpose(1, 2)
q, k = apply_rope(q, positions), apply_rope(k, positions)
# PyTorch 2.2+ 的 sdpa 会自动选 Flash-Attention v2/v3 后端
out: Float[Tensor, "B H S d_h"] = sdpa(q, k, v, is_causal=True)
out = out.transpose(1, 2).reshape(B, S, D)
return self.W_o(out)
3.2 每一步的维度与 FLOPs¶
| 步骤 | 输入 | 输出 | FLOPs | 备注 |
|---|---|---|---|---|
| W_qkv 投影 | [B,S,D] |
[B,S,3D] |
6 B S D² |
融合 QKV = 3 次独立投影 |
| 拆多头 reshape | [B,S,D] |
[B,H,S,d_h] |
0 | memory layout 变换 |
| RoPE | [B,H,S,d_h] |
[B,H,S,d_h] |
O(B H S d_h) |
负担可忽略 |
| Q K^T | — | [B,H,S,S] |
2 B H S² d_h |
O(S²) 瓶颈 |
| softmax | — | [B,H,S,S] |
O(B H S²) |
Flash-Attention 把它 fuse 掉了 |
| Attn · V | — | [B,H,S,d_h] |
2 B H S² d_h |
第二个 O(S²) |
| reshape | [B,H,S,d_h] |
[B,S,D] |
0 | |
| W_o | [B,S,D] |
[B,S,D] |
2 B S D² |
输出投影 |
Attention 总 FLOPs ≈ 8 B S D² + 4 B H S² d_h = 8 B S D² + 4 B S² D
- 短序列(
S ≪ D):投影主导,复杂度 O(S D²) - 长序列(
S ≫ D):QK/AV 主导,复杂度 O(S² D)
这就是上下文长度扩展到百万 token 时必须上 sparse/linear attention 的根本原因(详见 A2-2)。
四、FFN / SwiGLU¶
4.1 SwiGLU vs 标准 FFN¶
| 结构 | 公式 | d_ff 约定 | 参数 | 经验效果 |
|---|---|---|---|---|
| 标准 FFN | W_down(GELU(W_up(x))) |
4D |
8 D² |
原始 GPT-2 |
| SwiGLU | W_down(SiLU(W_gate(x)) ⊙ W_up(x)) |
8/3 D(对齐 256) |
8 D² |
同参数下 loss 低约 0.5% (PaLM / LLaMA) |
为什么 d_ff = 8/3 D:SwiGLU 多了一个 W_gate,要保持总参数量 3·d_ff·D = 8D² 不变,就得 d_ff = 8D/3。实现时取最近的 256 倍数(GPU 对齐)。
import torch.nn.functional as F
class SwiGLU_FFN(nn.Module):
def __init__(self, D: int, F_: int):
super().__init__()
self.W_gate = nn.Linear(D, F_, bias=False)
self.W_up = nn.Linear(D, F_, bias=False)
self.W_down = nn.Linear(F_, D, bias=False)
def forward(self, x: Float[Tensor, "B S D"]) -> Float[Tensor, "B S D"]:
gate: Float[Tensor, "B S F"] = F.silu(self.W_gate(x))
up: Float[Tensor, "B S F"] = self.W_up(x)
return self.W_down(gate * up)
4.2 FFN 的计算量¶
- 参数:
3 · D · F ≈ 8 D²(取F = 8D/3) - 每 token FLOPs:
2 · (3 · D · F) ≈ 16 D² - FFN FLOPs ≈ Attention 投影部分的 2 倍,所以 Transformer 参数大头在 FFN 而非 Attention
五、RMSNorm¶
去掉了 LayerNorm 的均值中心化 + bias,只做方差归一 + 可学习缩放。计算少 7% 左右,经验上精度持平。
class RMSNorm(nn.Module):
def __init__(self, D: int, eps: float = 1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(D))
self.eps = eps
def forward(self, x: Float[Tensor, "... D"]) -> Float[Tensor, "... D"]:
rms = x.pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
return x * rms * self.weight
| 维度 | LayerNorm | RMSNorm |
|---|---|---|
| 公式 | (x − μ) / σ · γ + β |
x / rms(x) · γ |
| 中心化 | ✅ 减均值 | ❌ |
| 可学习 bias | ✅ β | ❌ |
| 参数 | 2D |
D |
| 相对速度 | 1.0 | ~1.07× |
| 代表模型 | BERT、GPT-2 | LLaMA、T5、Qwen、Mistral |
六、参数量完整公式¶
对 D = d_model, H = n_heads, F = d_ff, L = n_layers, V = vocab_size:
| 组件 | 参数量 |
|---|---|
| Embedding | V · D |
| 每 Block · Attention | 4 D² (W_q/k/v/o, 忽略 bias) |
| 每 Block · FFN (SwiGLU) | 3 · D · F ≈ 8 D² (取 F=8D/3) |
| 每 Block · 2 × RMSNorm | 2 D |
| 总 Block 参数 | ≈ 12 D² per layer |
| 最终 RMSNorm | D |
| LM Head | V · D (通常与 Embedding 共享 → 0) |
近似公式:总参数 ≈ 12 L D² + 2 V D(embedding tied 时 LM Head 不额外算)。
速查表(LLaMA 系列)¶
| 模型 | L | D | H | F | V | 参数(12LD²) | 实际 |
|---|---|---|---|---|---|---|---|
| LLaMA-7B | 32 | 4096 | 32 | 11008 | 32000 | 6.44 B | 6.74 B |
| LLaMA-13B | 40 | 5120 | 40 | 13824 | 32000 | 12.58 B | 13.0 B |
| LLaMA-70B | 80 | 8192 | 64 | 28672 | 32000 | 64.42 B | 69.0 B |
公式和实际差 ~5%,来源:embedding(+V·D)、bias/norm 的小参数、GQA 带来的 KV 头折扣(A2-1)。
七、FLOPs 与训练预算¶
7.1 前向/反向 近似公式¶
- 前向每 token FLOPs ≈
2 N(N = 参数量)—— 每个参数参与一次乘加 = 2 FLOPs - 反向约是前向的 2 倍
- 训练总 FLOPs ≈ 6 × N × tokens(Chinchilla 论文里的规模律 baseline)
7.2 工程估算:LLaMA-7B 训 1T tokens¶
| 项 | 值 |
|---|---|
| 参数 N | 7 × 10⁹ |
| tokens | 1 × 10¹² |
| 总 FLOPs | 6 × 7e9 × 1e12 = 4.2 × 10²² FLOPs |
| H100 理论 BF16 吞吐 | 989 TFLOPs/s |
| 实际 MFU(混合并行) | 40–50% |
| 单 H100 实际吞吐 | ~450 TFLOPs/s |
| 单卡耗时 | 4.2e22 / 4.5e14 ≈ 2.6 × 10⁸ 秒 ≈ 3000 天 |
| 512 H100 实际耗时 | ~5.8 天 |
八、工程细节¶
8.1 QKV 投影融合¶
单次 Linear(D, 3D) 比三个独立 Linear(D, D) 快:
- 减少 kernel launch 次数(3 → 1)
- 更好的 GEMM 形状(更大的 M·N·K 对 tensor core 利用率友好)
- activation 只读一次
所有现代实现(Flash-Attention、vLLM、nanoGPT)都融合。GQA 场景下会融合成 Linear(D, D + 2·D_kv)(见 A2-1)。
8.2 Flash-Attention 的贡献¶
朴素 attention 把 [B,H,S,S] 的 score 矩阵写进 HBM,S=4096 时就是 2 GB 的 activation 内存。Flash-Attention 用 tiled online softmax 把整个 Q → K^T → softmax → V fuse 进一个 CUDA kernel,score 矩阵从来不落 HBM:
- 显存:O(S²) → O(S)
- 速度:2–4× 提升,序列越长越明显
- Flash-Attention 3(Hopper,2024)进一步用 WGMMA 异步指令 + FP8 → H100 上接近 75% 理论峰值
详见 FlashAttention-3 (2024) 和 C4 章。
8.3 参考实现¶
| 实现 | 特点 |
|---|---|
| nanoGPT | ~300 行 PyTorch,读懂 forward/backward 的最佳起点 |
| llama2.c | 纯 C 推理,理解 weight layout 和 inference-only 的极简实现 |
| gpt-fast | PyTorch 2.x 原生,torch.compile + int8/int4 量化参考 |
| torchtune | 官方 fine-tuning 栈,架构和训练分离解耦得很好 |
九、延伸问题¶
- KV Cache 的实现与显存估算 → A2(注意力机制全景)
- RoPE 的数学推导与外推方法 → A4-1、A4-2
- MoE 怎么把 FFN 拆成多个专家 → A3
- 权重初始化(Kaiming / μP / Parabolic Fitting) → D1
- 分布式切分:TP / PP / CP 下 Attention 怎么拆 → C1
- BF16 / FP8 训练下哪些层能量化 → D2
参考¶
- Vaswani et al. Attention Is All You Need (2017)
- Shazeer. GLU Variants Improve Transformer (2020) — SwiGLU 起源
- Zhang & Sennrich. Root Mean Square Layer Normalization (2019)
- Xiong et al. On Layer Normalization in the Transformer (2020) — Pre-Norm 理论
- Dao. FlashAttention-3 (2024)
- Raschka. The Big LLM Architecture Comparison (2024)
- Touvron et al. LLaMA 2: Open Foundation and Fine-Tuned Chat Models (2023)
一、完整前向传播¶
1.1 整体结构¶
一个 Decoder-Only Transformer(如 LLaMA)的前向传播:
class TransformerLM:
def __init__(self, vocab_size, d_model, n_layers, n_heads, d_ff):
self.embed = Embedding(vocab_size, d_model) # token → 向量
self.layers = [TransformerBlock(d_model, n_heads, d_ff) for _ in range(n_layers)]
self.norm = RMSNorm(d_model) # 最终归一化
self.lm_head = Linear(d_model, vocab_size, bias=False) # 向量 → logits
def forward(self, token_ids):
# token_ids: [batch, seq_len] — 整数
x = self.embed(token_ids) # [batch, seq_len, d_model]
for layer in self.layers:
x = layer(x) # [batch, seq_len, d_model]
x = self.norm(x) # [batch, seq_len, d_model]
logits = self.lm_head(x) # [batch, seq_len, vocab_size]
return logits
1.2 单个 Transformer Block¶
class TransformerBlock:
def __init__(self, d_model, n_heads, d_ff):
self.attn_norm = RMSNorm(d_model)
self.attn = CausalSelfAttention(d_model, n_heads)
self.ffn_norm = RMSNorm(d_model)
self.ffn = SwiGLU_FFN(d_model, d_ff)
def forward(self, x):
# Pre-Norm + Residual (不是 Post-Norm!)
x = x + self.attn(self.attn_norm(x)) # 残差 + 注意力
x = x + self.ffn(self.ffn_norm(x)) # 残差 + FFN
return x
1.3 Pre-Norm vs Post-Norm¶
| 对比 | Pre-Norm | Post-Norm |
|---|---|---|
| 公式 | x + Attn(Norm(x)) | Norm(x + Attn(x)) |
| 训练稳定性 | 更稳定 | 容易梯度爆炸 |
| 最终性能 | 略低(有争议) | 理论上更好但难训 |
| 使用情况 | LLaMA/Qwen/DeepSeek 全部使用 | GPT-2 使用,现已弃用 |
| 为什么? | 残差路径不经过 Norm → 梯度直通 | Norm 在残差路径上 → 梯度衰减 |
五、完整参数量计算¶
对于 d_model=D, n_heads=H, d_ff=F, n_layers=L, vocab_size=V:
| 组件 | 参数量公式 | LLaMA-7B 实际 | 为什么这么设计 |
|---|---|---|---|
| Embedding | V × D | 128256 × 4096 ≈ 525M | 每个 token 需要一个 D 维向量表示;vocab 越大表示能力越强但参数越多 |
| 每层 QKV 投影 | 3 × D² (或 GQA 时更少) | 3 × 4096² = 50M | Q/K/V 各需独立投影矩阵;GQA 可将 K/V 头数减少到 Q 的 ⅛,节省约 40% 参数 |
| 每层输出投影 | D² | 4096² = 17M | 将多头拼接结果映射回 d_model 空间,是多头信息融合的关键 |
| 每层 SwiGLU | 3 × D × F | 3 × 4096 × 11008 = 135M | FFN 是单层最大参数块(占 ~67%);d_ff=11008 是 8/3×4096 后做 256 对齐的结果 |
| 每层 RMSNorm ×2 | 2 × D | 2 × 4096 = 8K | Norm 参数极少但对训练稳定性至关重要;每层 attn 和 ffn 前各一个 |
| 最终 RMSNorm | D | 4096 | 最后一层输出到 LM Head 前做归一化,稳定 logits 的数值范围 |
| LM Head | D × V (通常与 Embedding 共享) | 共享 → 0 | 权重共享(weight tying)省 525M 参数且让输入输出在同一语义空间,小模型必用 |
| 每层总计 | ≈ 4D² + 3DF | ≈ 202M | Attention 占 ⅓,FFN 占 ⅔——这就是为什么 FFN 是参数效率优化的重点 |
| 组件 | 参数量公式 | LLaMA-7B 实际 | 为什么这么设计 |
|---|---|---|---|
| L 层总计 | L × (4D² + 3DF) | 32 × 202M = 6.46B | 层数 L 和宽度 D 的取舍:同等参数量下更深更窄通常优于更浅更宽 |
| 总计 | V×D + L×(4D²+3DF) | ≈ 6.98B | Embedding 占总参数 ~7.5%,但常被 Scaling Law 忽略(不算在"有效参数"内) |
实践笔记:面试常问"7B 模型的参数怎么算",记住近似公式 12LD² 即可快速估算(忽略 Embedding,假设 F≈8/3×D)。
近似公式¶
总参数 ≈ 12 × L × D² (当 F ≈ 8/3 × D 且忽略 embedding 时)
六、FLOPS 估算¶
训练一个 token 的前向+反向传播总 FLOPs ≈ 6 × N (N = 参数量)
| 推导 | 值 | 为什么 |
|---|---|---|
| 前向 FLOPs ≈ 2N | 每个参数一次乘法一次加法 | 矩阵乘法 [m,k]@[k,n] 的 FLOPs = 2mkn,而参数量就是 k×n,所以每个参数贡献 2 FLOPs |
| 反向 FLOPs ≈ 4N | 约为前向的 2 倍 | 反向需要对输入和权重各做一次矩阵乘来计算梯度,共 2 次前向等量计算 |
| 总计 ≈ 6N per token | 这就是 Chinchilla 计算公式 C ≈ 6ND 的来源 | 6N 是一阶近似,实际包含 Norm/Softmax 等开销约多 5-10%,但数量级准确 |
训练 LLaMA-7B 的总 FLOPs:6 × 7B × 1T tokens ≈ 4.2 × 10²² FLOPs
H100 需要多久:4.2e22 / (990e12 × 0.45 MFU) ≈ 94,000 GPU-hours ≈ 512 H100 跑 7.7 天
七、追问:这些知识够复现一个 Transformer 了吗?¶
还差:
-
KV Cache 的完整实现(推理用)→ 见 A2
-
RoPE 的数学推导(为什么旋转能编码位置)→ 见 A4
-
如何初始化权重(Xavier? Kaiming? 还是别的?)→ 见 D3
-
训练时的 loss 计算和梯度 → 见 D3
-
分布式切分时 Attention 怎么拆 → 见 C1
每个延伸问题在对应子文档中深入。
参考链接¶
↑ 上级 · A. 基础理论