跳转至

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

参考

一、完整前向传播

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% 参数
每层输出投影 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. 基础理论