Transformer 深度详解 — 从公式到代码的完整推导
#Transformer · #自注意力 · #多头注意力 · #位置编码 · #残差连接 · #LayerNorm · #编码器 · #解码器
本文是 Transformer 架构的完整深度解读,从数学公式、直观理解到完整代码实现。承接 人工智能 & 深度学习 的基础内容,建议先阅读总览再进入本文。
历史背景:为什么需要 Transformer?
mermaid
timeline
title NLP 序列建模技术的三次范式转移
2014以前
: RNN/LSTM 时代 — 序列递推, 无法并行, 长距离依赖困难
2014-2015
: Attention 诞生 — Bahdanau/Luong Attention 解决了信息瓶颈
2017.06
: Transformer 诞生 — "Attention is All You Need" 抛弃循环, 完全并行
2018至今
: 预训练时代 — BERT, GPT, T5 在此基础上发展出 LLM在 Transformer 之前,Seq2Seq 模型的工作方式是:
- Encoder RNN 把整个输入压缩成一个固定长度的向量
- Decoder RNN 从这个向量中"解压"出目标序列
问题很明显:不论输入多长,都必须塞进一个固定大小的向量。就像让你用一句话总结整本书——必然丢失大量信息。
Attention 机制提供了一个优雅的解决方案:Decoder 在每个时刻都可以"回顾"Encoder 的所有隐藏状态,选出最相关的部分。但 Attention 本身仍然是"附加"在 RNN 上的,串行计算的根本瓶颈没有解决。
Transformer 的大胆创新:如果 Attention 这么好用,为什么还要保留 RNN?
Transformer 整体架构
mermaid
graph TB
subgraph "输入处理"
SRC["源序列<br/>'我 爱 AI'"] --> TOK["Tokenization<br/>['我', '爱', 'AI']"]
TOK --> EMB["词嵌入<br/>d_model = 512"]
EMB --> PE["+ 位置编码<br/>sin/cos 编码"]
end
subgraph "编码器 × N (N=6)"
PE --> ENC1["编码器层 1"]
ENC1 --> ENC2["编码器层 2"]
ENC2 --> ENCN["..."]
ENCN --> ENC6["编码器层 6"]
end
subgraph "解码器 × N (N=6)"
DEC_EMB["目标序列嵌入 + 位置编码"] --> DEC1["解码器层 1"]
ENC6 --> DEC1
DEC1 --> DEC2["解码器层 2"]
DEC2 --> DECN["..."]
DECN --> DEC6["解码器层 6"]
DEC6 --> LINEAR["线性层"]
LINEAR --> SOFTMAX["Softmax"]
SOFTMAX --> OUTPUT["输出概率分布"]
end
style PE fill:#3498db,color:#fff
style ENC6 fill:#e74c3c,color:#fff
style DEC6 fill:#2ecc71,color:#fffGPT 只用了 Decoder,BERT 只用了 Encoder。 分别代表了"生成式"和"理解式"两条路线。
1. 输入表示层
1.1 Tokenization — 把文本变成数字
Transformer 不能直接处理文字,需要先 Tokenization:
python
# 示例:不同 Tokenization 策略
text = "I love artificial intelligence."
# 1. 词级分词 (Word-level)
tokens_word = ["I", "love", "artificial", "intelligence", "."]
# 2. 子词分词 (Subword) — BPE/WordPiece/SentencePiece
# BPE: 频率最高的字符对合并,生僻词分解为子词
tokens_bpe = ["I", "love", "artificial", "intel", "ligence", "."]
# 3. 字符级分词 (Character-level)
tokens_char = ["I", "l", "o", "v", "e", " ", "a", ...]现代 LLM 普遍使用 BPE(Byte Pair Encoding)或 SentencePiece。优势:
- 开放的词汇表(没有 OOV 问题)
- 高频词保持完整,低频词分解为子词
- 中英文混合处理友好
python
from transformers import AutoTokenizer
# GPT-2 的 Tokenizer(BPE)
tokenizer = AutoTokenizer.from_pretrained("gpt2")
tokens = tokenizer("I love artificial intelligence.")
print(f"Token IDs: {tokens['input_ids']}")
print(f"Tokens: {[tokenizer.decode([t]) for t in tokens['input_ids']]}")
# 输出类似: ['I', 'Ġlove', 'Ġart', 'ificial', 'Ġint', 'elligence', '.']
# 'Ġ' 表示前面的空格BPE 训练算法详解
BPE 的核心操作是:每轮统计所有相邻 Token 对的频率,合并频率最高的那一对,直到达到目标词表大小。
python
"""
BPE 训练算法的完整实现
"""
from collections import Counter, defaultdict
from typing import List, Dict, Tuple
def train_bpe(corpus: List[str], vocab_size: int) -> Dict[Tuple[str, str], str]:
"""
从头训练 BPE 词表
Args:
corpus: 训练语料,每段文本以字符列表表示
vocab_size: 目标词表大小(含基础字符)
Returns:
merge_rules: {(a, b): ab} 合并规则(按顺序应用)
示例:
>>> corpus = [['l', 'o', 'w', ' '], ['l', 'o', 'w', 'e', 'r', ' ']]
>>> rules = train_bpe(corpus, 10)
>>> rules
{('l', 'o'): 'lo', ('lo', 'w'): 'low'}
"""
# 1. 初始化:把每个文本拆成字符序列(加结束符</w>)
vocab = defaultdict(int)
for text in corpus:
tokens = list(text) + ["</w>"]
vocab[tuple(tokens)] += 1
# 2. 收集基础字符作为初始词表
char_set = set()
for text in corpus:
char_set.update(text)
char_set.add("</w>")
vocab_list = sorted(char_set) # 初始词表
# 3. 迭代合并
merge_rules = {}
num_merges = vocab_size - len(vocab_list)
for step in range(num_merges):
# 3a. 统计所有相邻 token 对的频率
pair_counts = Counter()
for token_seq, count in vocab.items():
for i in range(len(token_seq) - 1):
pair = (token_seq[i], token_seq[i + 1])
pair_counts[pair] += count
if not pair_counts:
break
# 3b. 选择频率最高的 pair
best_pair = max(pair_counts, key=pair_counts.get)
a, b = best_pair
new_token = a + b.replace("</w>", "")
merge_rules[best_pair] = new_token
vocab_list.append(new_token)
# 3c. 应用合并规则到整个词表
new_vocab = defaultdict(int)
for token_seq, count in vocab.items():
new_seq = list(token_seq)
i = 0
while i < len(new_seq) - 1:
if (new_seq[i], new_seq[i + 1]) == best_pair:
new_seq[i] = new_seq[i] + new_seq[i + 1]
new_seq.pop(i + 1)
i += 1
new_vocab[tuple(new_seq)] += count
vocab = new_vocab
print(f"Step {step+1}: merged '{a}' + '{b}' → '{new_token}' "
f"(freq={pair_counts[best_pair]})")
return merge_rules
# 演示: 小语料上训练 BPE
demo_corpus = [
"low low low lower lowest",
]
# 按空格分词再拆字符
words = demo_corpus[0].split()
char_level = [" ".join(list(w)) for w in words]
print("训练语料:", words)
rules = train_bpe(
corpus=[list(w) + ["</w>"] for w in words],
vocab_size=15,
)BPE 的编码过程
训练完成后,用合并规则对新文本做分词:
python
def encode_bpe(text: str, merge_rules: Dict[Tuple[str, str], str]) -> List[str]:
"""
用已训练的 BPE 规则对新文本编码
原理:重复应用合并规则直到无法再合并
"""
tokens = list(text) + ["</w>"]
while True:
# 找当前序列中优先级最高的可合并 pair
# (规则按训练顺序排列,先训练的优先)
best_pair = None
best_pos = -1
best_priority = float("inf")
for i in range(len(tokens) - 1):
pair = (tokens[i], tokens[i + 1])
if pair in merge_rules:
priority = list(merge_rules.keys()).index(pair)
if priority < best_priority:
best_pair = pair
best_pos = i
best_priority = priority
if best_pair is None:
break
# 合并
tokens[best_pos] = merge_rules[best_pair]
tokens.pop(best_pos + 1)
return tokens
# 编码示例
print(encode_bpe("lower", rules))
# 输出: ['low', 'er', '</w>']BPE 的本质:贪心地将语料中最常共现的字符对压缩为一个新 token。训练得到的合并规则就是词表,数量(vocab_size)是人为设定的超参数。GPT-2 用 50K,LLaMA 用 32K,Qwen 用 152K。
1.2 词嵌入 (Input Embedding)
将离散的 Token ID 映射为稠密向量:
其中
python
class InputEmbedding(nn.Module):
"""输入嵌入层 — 将 Token ID 转为稠密向量"""
def __init__(self, vocab_size: int, d_model: int):
super().__init__()
self.embedding = nn.Embedding(vocab_size, d_model)
self.d_model = d_model
# 初始化策略很重要!太大会导致 Softmax 过饱和
nn.init.normal_(self.embedding.weight, mean=0, std=1.0)
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: (batch, seq_len) → (batch, seq_len, d_model)
# 乘以 √d_model 是为了让 embedding 的方差稳定
return self.embedding(x) * math.sqrt(self.d_model)为什么乘以 √d_model? Embedding 的初始方差约为 1,与位置编码相加后如果不缩放,嵌入的值相对位置编码太小,导致位置信息主导。乘以 √d_model 保持方差平衡。
2. 位置编码 (Positional Encoding)
为什么需要位置编码?
Attention 机制是置换不变的:交换输入中两个词的位置,Attention 输出也相应交换。也就是说,如果不加位置信息,"A 追 B" 和 "B 追 A" 对模型来说是一样的。
原论文使用正弦/余弦位置编码:
其中:
:词在序列中的位置(0, 1, 2, ...) :维度索引(0, 1, ..., d_model/2 - 1) :模型维度
关键性质:
- 每个位置的编码是唯一的
- 相邻位置的编码相似(位置 5 和 6 的编码比 5 和 50 更相似)
- 不同维度对应不同频率——低维度编码"粗粒度"位置,高维度编码"细粒度"位置
可以由 的线性函数表示 → 模型能学到相对位置关系
python
class SinusoidalPositionalEncoding(nn.Module):
"""正弦/余弦位置编码 — 原论文实现"""
def __init__(self, d_model: int, max_len: int = 5000, dropout: float = 0.1):
super().__init__()
self.dropout = nn.Dropout(p=dropout)
# 创建位置编码矩阵 (max_len, d_model)
pe = torch.zeros(max_len, d_model)
# pos: (max_len, 1)
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
# div_term: (d_model/2,) — 不同维度的频率
# 10000^(2i/d_model) 其中 i 从 0 到 d_model/2-1
div_term = torch.exp(
torch.arange(0, d_model, 2).float() *
(-math.log(10000.0) / d_model)
)
# 偶数维度: sin(pos * div_term)
pe[:, 0::2] = torch.sin(position * div_term)
# 奇数维度: cos(pos * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
# 增加 batch 维度: (1, max_len, d_model)
pe = pe.unsqueeze(0)
# 注册为 buffer(不参与训练,但会随模型保存)
self.register_buffer('pe', pe)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
x: (batch, seq_len, d_model)
返回: x + 位置编码
"""
# 取对应长度: (1, seq_len, d_model)
seq_len = x.size(1)
x = x + self.pe[:, :seq_len, :]
return self.dropout(x)位置编码可视化
不同位置(纵轴)在不同维度(横轴)上的编码值。深色=大正值,浅色=大负值:
python
import matplotlib.pyplot as plt
import numpy as np
def plot_positional_encoding():
"""绘制位置编码的热力图展示其频率特性"""
pe = SinusoidalPositionalEncoding(d_model=128, max_len=100)
encoding = pe.pe.squeeze(0).numpy() # (100, 128)
fig, axes = plt.subplots(1, 2, figsize=(14, 5))
# 左: 热力图 (横轴=维度, 纵轴=位置)
im = axes[0].imshow(encoding, aspect='auto', cmap='RdBu')
axes[0].set_xlabel('维度')
axes[0].set_ylabel('位置')
axes[0].set_title('位置编码热力图\n(低维度 = 低频, 高维度 = 高频)')
plt.colorbar(im, ax=axes[0])
# 右: 不同位置的余弦相似度
pos_0 = encoding[0] # 位置 0 的编码
similarities = np.array([
np.dot(pos_0, encoding[i]) / (np.linalg.norm(pos_0) * np.linalg.norm(encoding[i]))
for i in range(100)
])
axes[1].plot(similarities)
axes[1].set_xlabel('位置差')
axes[1].set_ylabel('与位置0的余弦相似度')
axes[1].set_title('相邻位置编码的相关性')
axes[1].axhline(y=0, color='gray', linestyle='--')
plt.tight_layout()
plt.savefig('positional_encoding.png', dpi=150)
# plt.show()
# plot_positional_encoding()其他位置编码方案
| 方案 | 说明 | 使用模型 |
|---|---|---|
| Sinusoidal (原版) | sin/cos 函数,固定不学习 | 原版 Transformer |
| Learned | 可学习的 Embedding | BERT, GPT, ViT |
| RoPE (旋转位置编码) | 通过旋转矩阵编码相对位置 | LLaMA, Qwen, Mistral |
| ALiBi | 在 Attention 分数上加线性偏置 | BLOOM |
| 相对位置编码 | 编码词对之间的相对距离 | Transformer-XL |
RoPE 是当前 LLM 的首选。它将位置信息编码为旋转操作,天然处理相对位置,支持任意长度的外推。
RoPE 的核心思想:
即旋转后的 Q 和 K 的点积只依赖于它们的相对位置
3. 自注意力机制 — Transformer 的核心
3.1 从直觉出发:什么是"注意力"?
在正式公式之前,先用一个例子理解什么是"查询-键-值"(Query-Key-Value):
想象你正在阅读下面这句话:
"猫追老鼠因为它饿了"
你的大脑需要判断"它"指代什么。大脑怎么做?
- 生成一个 查询(Query):"谁有可能被'它'指代?" → 查询包含代词的特征(第三人称、单数、动物)
- 句子中每个词都有一个 键(Key):描述"我是什么类型的词"
- "猫"的 Key:
[动物, 主语, 第三人称] - "追"的 Key:
[动作, 谓语, ...] - "老鼠"的 Key:
[动物, 宾语, 第三人称]
- "猫"的 Key:
- 查询与键匹配:Q("它") 与 K("猫") 的点积 = 相似度分数。谁的分数高就说明 Q 更"关注"谁
- 每个词还有一个 值(Value):表示"如果我被选中,我要提供什么信息"
- 加权求和:用匹配分数做权重,对 Value 加权平均 → 得到最终输出
公式表达的就是这个"查字典"的过程:
3.2 分步拆解:一步一步看懂公式
下面用一句 3 个词的话 "我 爱 AI" 来逐步走完整个计算。假设每个词已经变成了一个 4 维向量(实际模型是 512 或更多维):
d_k = 4(每头的维度,实际模型为 64)
Q = [[1.0, 0.5, 0.2, 0.8], ← "我"的查询向量
[0.3, 1.2, 0.7, 0.1], ← "爱"的查询向量
[0.6, 0.3, 1.1, 0.4]] ← "AI"的查询向量 → 形状 (3, 4)
K = [[0.8, 0.3, 0.5, 0.2], ← "我"的键向量
[0.4, 0.9, 0.1, 0.7], ← "爱"的键向量
[0.7, 0.2, 0.8, 0.3]] ← "AI"的键向量 → 形状 (3, 4)
V = [[0.1, 0.4, 0.7, 0.2], ← "我"的值向量
[0.5, 0.2, 0.8, 0.3], ← "爱"的值向量
[0.9, 0.6, 0.1, 0.5]] ← "AI"的值向量 → 形状 (3, 4)第 1 步:Q × Kᵀ → 计算"谁和谁相关"
QKᵀ = Q @ K.T,形状:(3,4) × (4,3) → (3,3)
Q[0]·K[0] = 1.0×0.8 + 0.5×0.3 + 0.2×0.5 + 0.8×0.2 = 0.80+0.15+0.10+0.16 = 1.21 ← "我"和"我"的相关性
Q[0]·K[1] = 1.0×0.4 + 0.5×0.9 + 0.2×0.1 + 0.8×0.7 = 0.40+0.45+0.02+0.56 = 1.43 ← "我"和"爱"的相关性
Q[0]·K[2] = 1.0×0.7 + 0.5×0.2 + 0.2×0.8 + 0.8×0.3 = 0.70+0.10+0.16+0.24 = 1.20 ← "我"和"AI"的相关性
未缩放的注意力分数(Scores):
"我" "爱" "AI"
"我" [1.21, 1.43, 1.20] ← 第一行:"我"的Q 对每个词 K 的匹配分数
"爱" [0.93, 1.89, 1.12] ← 第二行:"爱"的Q 对每个词 K 的匹配分数
"AI" [1.04, 1.30, 1.47] ← 第三行:"AI"的Q 对每个词 K 的匹配分数第 2 步:÷ √dₖ → 防止分数过大
d_k = 4,√d_k = 2
缩放后:
"我" "爱" "AI"
"我" [0.61, 0.72, 0.60]
"爱" [0.47, 0.95, 0.56]
"AI" [0.52, 0.65, 0.74]为什么要除 √dₖ? Q 和 K 的元素是独立随机变量(均值为 0,方差为 1)。dₖ 个元素相乘相加,点积结果的方差会涨到 dₖ。当 dₖ=64 时,点积可能达到 ±√64=±8。这么大的值丢进 softmax 会变成极端分布(一个接近 1,其余接近 0),梯度几乎为零 → 模型学不动。除以 √dₖ 把方差拉回 1。
第 3 步:softmax → 把分数变成"权重"(∑=1)
softmax 对每一行独立计算。以第一行 [0.61, 0.72, 0.60] 为例:
e^0.61 = 1.84, e^0.72 = 2.05, e^0.60 = 1.82
总和 = 1.84 + 2.05 + 1.82 = 5.71
权重 = [1.84/5.71, 2.05/5.71, 1.82/5.71] = [0.322, 0.359, 0.319]
完整注意力权重矩阵:
"我" "爱" "AI"
"我" [0.322, 0.359, 0.319] ← "我"关注每个词的比例("爱"略高)
"爱" [0.258, 0.415, 0.327] ← "爱"最关注自己(0.415)
"AI" [0.281, 0.320, 0.399] ← "AI"最关注自己(0.399)softmax 做了什么? 把每一行的数值转化为概率分布——值大的放大,值小的压扁。
[0.61, 0.72, 0.60]三个数本来差不多,softmax 后还差不多(各 ~1/3)。但如果有一个远大于其他几个(如[2.0, 0.1, 0.1]),softmax 后那个会吞掉几乎所有权重。
第 4 步:权重 × V → 加权平均("融合信息")
Output = 权重矩阵 @ V
第一行("我"位置的输出):
= 0.322 × V["我"] + 0.359 × V["爱"] + 0.319 × V["AI"]
= 0.322×[0.1,0.4,0.7,0.2] + 0.359×[0.5,0.2,0.8,0.3] + 0.319×[0.9,0.6,0.1,0.5]
= [0.032,0.129,0.225,0.064] + [0.180,0.072,0.287,0.108] + [0.287,0.191,0.032,0.160]
= [0.499, 0.392, 0.544, 0.332]
最终输出:
"我" → [0.499, 0.392, 0.544, 0.332] ← 融入了所有词的语义
"爱" → [0.416, 0.391, 0.586, 0.341]
"AI" → [0.459, 0.428, 0.529, 0.364]3.3 公式总结 — 一张图看懂
mermaid
graph TD
subgraph "缩放点积注意力 — 四步计算"
Q["Q (Query)<br/>我想找什么?"] --> DOT["1️⃣ 计算相关性<br/>S = Q·Kᵀ<br/>(batch, n_heads, seq, seq)"]
K["K (Key)<br/>我能提供什么?"] --> DOT
DOT --> SCALE["2️⃣ 缩放<br/>S / √dₖ<br/>防止点积过大"]
SCALE --> MASK["(可选) 掩码<br/>因果掩码/padding掩码"]
MASK --> SOFTMAX["3️⃣ Softmax 归一化<br/>A = softmax(S/√dₖ)"]
V["V (Value)<br/>我实际有什么?"] --> WEIGHTED["4️⃣ 加权求和<br/>Output = A·V"]
SOFTMAX --> WEIGHTED
WEIGHTED --> OUT["输出<br/>(batch, n_heads, seq, d_k)"]
end
style SOFTMAX fill:#e74c3c,color:#fff
style WEIGHTED fill:#2ecc71,color:#fff公式:
| 步骤 | 操作 | 形状变化 | 直觉 |
|---|---|---|---|
| 1 | (seq, d_k) × (d_k, seq) → (seq, seq) | 每个词和每个词"聊天",看谁和谁相关 | |
| 2 | (seq, seq) | 控制数值大小,别让 softmax 饱和 | |
| 3 | (seq, seq) | 把分数变成"谁占多少权重" | |
| 4 | (seq, seq) × (seq, d_k) → (seq, d_k) | 按权重融合所有词的信息 |
python
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
class ScaledDotProductAttention(nn.Module):
"""缩放点积注意力 — Transformer 的原子操作"""
def __init__(self, dropout: float = 0.1):
super().__init__()
self.dropout = nn.Dropout(dropout)
def forward(self, Q: torch.Tensor, K: torch.Tensor, V: torch.Tensor,
mask: torch.Tensor = None) -> tuple:
"""
Args:
Q: (batch, n_heads, seq_len, d_k)
K: (batch, n_heads, seq_len, d_k)
V: (batch, n_heads, seq_len, d_k)
mask: (batch, 1, seq_len, seq_len) 或 (seq_len, seq_len)
Returns:
output: (batch, n_heads, seq_len, d_k)
attention_weights: (batch, n_heads, seq_len, seq_len)
"""
d_k = Q.size(-1)
# 1. 计算注意力分数: Q·Kᵀ (batch, n_heads, seq, seq)
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)
# 2. 掩码处理
if mask is not None:
# mask 中值为 0 的位置 → 注意力分数设为 -∞
scores = scores.masked_fill(mask == 0, float('-inf'))
# 3. Softmax 归一化
attention_weights = F.softmax(scores, dim=-1)
attention_weights = self.dropout(attention_weights)
# 4. 加权求和
output = torch.matmul(attention_weights, V)
return output, attention_weights3.3 因果掩码 (Causal Mask) — GPT 的关键
GPT 是自回归模型,生成第
python
def create_causal_mask(seq_len: int) -> torch.Tensor:
"""创建因果掩码(下三角矩阵)
结果:
[[1, 0, 0, 0],
[1, 1, 0, 0],
[1, 1, 1, 0],
[1, 1, 1, 1]]
位置 (i,j): i 能看到 j 当且仅当 j ≤ i
"""
mask = torch.tril(torch.ones(seq_len, seq_len))
return mask # 1=可见, 0=不可见(会变-inf)4. 多头注意力 (Multi-Head Attention)
4.1 为什么需要多头?
单头注意力的表达能力有限。多头意味着:
- 不同的"注意力头"关注不同类型的依赖关系
- 有的头关注语法关系(主语-谓语),有的关注语义关系(代词-指代对象)
- 每个头的维度更小(
),总计算量不变
mermaid
graph TD
X["输入 X<br/>(batch, seq, d_model)"] --> QP["Linear_Q: d_model → d_model"]
X --> KP["Linear_K: d_model → d_model"]
X --> VP["Linear_V: d_model → d_model"]
QP --> SPLIT_Q["拆分为 h 个头<br/>(batch, h, seq, d_k)"]
KP --> SPLIT_K["拆分为 h 个头<br/>(batch, h, seq, d_k)"]
VP --> SPLIT_V["拆分为 h 个头<br/>(batch, h, seq, d_k)"]
SPLIT_Q --> ATTN0["头 0: Attention₀"]
SPLIT_K --> ATTN0
SPLIT_V --> ATTN0
SPLIT_Q --> ATTN1["头 1: Attention₁"]
SPLIT_K --> ATTN1
SPLIT_V --> ATTN1
SPLIT_Q --> ATTN_DOTS["···"]
SPLIT_K --> ATTN_DOTS
SPLIT_V --> ATTN_DOTS
SPLIT_Q --> ATTN7["头 7: Attention₇"]
SPLIT_K --> ATTN7
SPLIT_V --> ATTN7
ATTN0 --> CONCAT["拼接所有头<br/>Concat(h₀, h₁, ..., h₇)"]
ATTN1 --> CONCAT
ATTN_DOTS --> CONCAT
ATTN7 --> CONCAT
CONCAT --> OUT_PROJ["Linear_O: d_model → d_model"]
OUT_PROJ --> Y["输出"]
style ATTN0 fill:#e74c3c,color:#fff
style ATTN1 fill:#3498db,color:#fff
style ATTN7 fill:#2ecc71,color:#fff公式:
其中
原论文取
python
class MultiHeadAttention(nn.Module):
"""完整的多头注意力实现"""
def __init__(self, d_model: int = 512, n_heads: int = 8, dropout: float = 0.1):
super().__init__()
assert d_model % n_heads == 0, f"d_model ({d_model}) 必须能被 n_heads ({n_heads}) 整除"
self.d_model = d_model
self.n_heads = n_heads
self.d_k = d_model // n_heads # 每个头的维度 = 512/8 = 64
self.d_v = d_model // n_heads
# Q, K, V 的线性投影 — 合并所有头的参数(效率更高)
self.W_q = nn.Linear(d_model, d_model) # (512 → 512)
self.W_k = nn.Linear(d_model, d_model)
self.W_v = nn.Linear(d_model, d_model)
# 输出投影
self.W_o = nn.Linear(d_model, d_model)
self.attention = ScaledDotProductAttention(dropout)
self.dropout = nn.Dropout(dropout)
def forward(self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor,
mask: torch.Tensor = None) -> torch.Tensor:
"""
Args:
query: (batch, seq_len, d_model)
key: (batch, seq_len, d_model)
value: (batch, seq_len, d_model)
mask: (batch, seq_len, seq_len) — 用于 padding/causal
Returns:
output: (batch, seq_len, d_model)
"""
batch_size = query.size(0)
# 1. 线性投影: d_model → d_model
Q = self.W_q(query) # (batch, seq, d_model)
K = self.W_k(key)
V = self.W_v(value)
# 2. 拆分为多头: (batch, seq, d_model) → (batch, n_heads, seq, d_k)
Q = Q.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
K = K.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
V = V.view(batch_size, -1, self.n_heads, self.d_v).transpose(1, 2)
# 3. 缩放点积注意力
# 输出: (batch, n_heads, seq, d_v), (batch, n_heads, seq, seq)
attn_out, _ = self.attention(Q, K, V, mask)
# 4. 合并多头: (batch, n_heads, seq, d_v) → (batch, seq, d_model)
attn_out = attn_out.transpose(1, 2).contiguous()
attn_out = attn_out.view(batch_size, -1, self.d_model)
# 5. 输出投影
output = self.W_o(attn_out)
return output
class MultiHeadCrossAttention(nn.Module):
"""交叉注意力 — Decoder 中用于关注 Encoder 输出
与自注意力的区别:
- 自注意力: Q=K=V (来自同一序列)
- 交叉注意力: Q 来自 Decoder, K,V 来自 Encoder
"""
def __init__(self, d_model: int = 512, n_heads: int = 8, dropout: float = 0.1):
super().__init__()
self.d_model = d_model
self.n_heads = n_heads
self.d_k = d_model // n_heads
# Q 来自 Decoder 自己的输出
self.W_q = nn.Linear(d_model, d_model)
# K, V 来自 Encoder 输出
self.W_k = nn.Linear(d_model, d_model)
self.W_v = nn.Linear(d_model, d_model)
self.W_o = nn.Linear(d_model, d_model)
self.attention = ScaledDotProductAttention(dropout)
def forward(self, decoder_hidden: torch.Tensor,
encoder_output: torch.Tensor,
mask: torch.Tensor = None) -> torch.Tensor:
"""
Args:
decoder_hidden: (batch, tgt_len, d_model) — Decoder 当前层的输出
encoder_output: (batch, src_len, d_model) — Encoder 最终输出
mask: (batch, tgt_len, src_len) — Encoder 的 padding mask
"""
batch_size = decoder_hidden.size(0)
# Q: Decoder 查询 Encoder
Q = self.W_q(decoder_hidden).view(
batch_size, -1, self.n_heads, self.d_k
).transpose(1, 2)
# K, V: Encoder 被查询的内容
K = self.W_k(encoder_output).view(
batch_size, -1, self.n_heads, self.d_k
).transpose(1, 2)
V = self.W_v(encoder_output).view(
batch_size, -1, self.n_heads, self.d_k
).transpose(1, 2)
attn_out, _ = self.attention(Q, K, V, mask)
attn_out = attn_out.transpose(1, 2).contiguous().view(
batch_size, -1, self.d_model
)
return self.W_o(attn_out)4.2 GQA/MQA — 现代 LLM 的注意力优化
随着模型变大,KV Cache 的内存消耗成为瓶颈。GQA (Grouped Query Attention) 和 MQA (Multi-Query Attention) 通过减少 KV 头的数量来节省内存:
| 方案 | Q 头数 | K/V 头数 | KV Cache 节省 | 质量损失 | 使用模型 |
|---|---|---|---|---|---|
| MHA | h | h | 基准 | 无 | 原版 Transformer, BERT |
| GQA | h | g (< h) | h/g 倍 | 几乎无 | LLaMA 2/3 (g=8), Mistral |
| MQA | h | 1 | h 倍 | 轻微 | PaLM, Gemini |
5. 前馈网络 (Feed-Forward Network)
5.1 标准 FFN
每个位置独立应用的两层全连接网络:
- 输入维度:
(512) - 中间维度:
(2048, 4×) - 输出维度:
(512)
python
class FeedForward(nn.Module):
"""位置级前馈网络 (Position-wise FFN)
标准实现: d_model → d_ff → d_model
"""
def __init__(self, d_model: int = 512, d_ff: int = 2048,
dropout: float = 0.1, activation: str = 'relu'):
super().__init__()
self.linear1 = nn.Linear(d_model, d_ff)
self.linear2 = nn.Linear(d_ff, d_model)
self.dropout = nn.Dropout(dropout)
# 激活函数选择
if activation == 'relu':
self.activation = nn.ReLU()
elif activation == 'gelu':
self.activation = nn.GELU()
elif activation == 'swiglu':
# SwiGLU 需要特殊的结构(两个线性变换+门控)
self.gate = nn.Linear(d_model, d_ff)
self.activation = nn.SiLU() # Swish/SiLU
else:
raise ValueError(f"Unknown activation: {activation}")
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
x: (batch, seq_len, d_model)
"""
if hasattr(self, 'gate'):
# SwiGLU: (xW_g ⊙ σ(xW_1)) · W_2
return self.linear2(
self.dropout(
self.activation(self.linear1(x)) * self.gate(x)
)
)
else:
# 标准 FFN: σ(xW_1 + b_1)W_2 + b_2
return self.linear2(
self.dropout(
self.activation(self.linear1(x))
)
)5.2 SwiGLU — 新一代 LLM 的标准
SwiGLU 已成为 LLaMA 2/3、PaLM 等模型的标配:
其中
6. 残差连接与层归一化
6.1 残差连接 (Residual Connection)
残差连接让梯度可以"直接流过"每一层,是训练深层网络的关键:
mermaid
graph LR
X["输入 x"] --> SUB["子层<br/>Attention / FFN"]
X --> ADD["⊕ 相加"]
SUB --> ADD
ADD --> LN["LayerNorm"]
LN --> OUT["输出"]
style ADD fill:#f39c12,color:#fff
style LN fill:#e74c3c,color:#fff为什么残差有效? 反向传播时,梯度 = 直接路径 + 子层路径。即使子层梯度消失,直接路径仍能传回梯度。
6.2 层归一化 (Layer Normalization)
与 Batch Normalization 按 batch 维度归一化不同,LayerNorm 按特征维度归一化:
其中:
(均值) (方差) 是可学习的缩放和偏移参数
python
class LayerNorm(nn.Module):
"""层归一化的手工实现 — 理解其原理"""
def __init__(self, d_model: int, eps: float = 1e-6):
super().__init__()
self.gamma = nn.Parameter(torch.ones(d_model)) # 可学习缩放
self.beta = nn.Parameter(torch.zeros(d_model)) # 可学习偏移
self.eps = eps
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: (batch, seq_len, d_model)
mean = x.mean(dim=-1, keepdim=True) # (batch, seq_len, 1)
std = x.std(dim=-1, keepdim=True) # (batch, seq_len, 1)
# 归一化 + 可学习参数
return self.gamma * (x - mean) / (std + self.eps) + self.beta为什么 Transformer 用 LayerNorm 而不是 BatchNorm?
- BatchNorm 依赖 batch 统计量,batch 太小不稳定
- 序列长度可变时,BatchNorm 需要处理 padding
- LayerNorm 在推理时不依赖 batch,行为完全确定
Pre-LN vs Post-LN
现代 LLM 普遍使用 Pre-LN(先 Norm 再子层),训练更稳定:
| 变体 | 公式 | 特点 |
|---|---|---|
| Post-LN (原版) | 原论文实现,需要 warmup | |
| Pre-LN (现代) | 训练更稳定,无需 warmup |
原论文使用 Post-LN:先残差相加,再 LayerNorm。流程是
x → Sublayer(x) → x + Sublayer(x) → LN(·) → out。现代 LLM 普遍用 Pre-LN:先对输入做 Norm,再做子层计算,再残差相加。流程是x → LN(x) → Sublayer(LN(x)) → x + Sublayer(LN(x)) → out。Pre-LN 让梯度流更平滑,省去了 warmup 阶段。
7. 完整的编码器层与解码器层
编码器层
python
class TransformerEncoderLayer(nn.Module):
"""Transformer 编码器层
结构: 自注意力 → 残差+Norm → FFN → 残差+Norm
"""
def __init__(self, d_model: int = 512, n_heads: int = 8,
d_ff: int = 2048, dropout: float = 0.1,
activation: str = 'relu'):
super().__init__()
# 多头自注意力(Q=K=V,都是编码器自己的输出)
self.self_attn = MultiHeadAttention(d_model, n_heads, dropout)
# 前馈网络
self.feed_forward = FeedForward(d_model, d_ff, dropout, activation)
# 两个 LayerNorm(每个子层后一个)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x: torch.Tensor, src_mask: torch.Tensor = None) -> torch.Tensor:
"""
Args:
x: (batch, src_len, d_model)
src_mask: padding mask (batch, 1, 1, src_len)
"""
# 子层1: 自注意力 + 残差 + Norm
attn_out = self.self_attn(x, x, x, src_mask)
x = self.norm1(x + self.dropout(attn_out))
# 子层2: FFN + 残差 + Norm
ffn_out = self.feed_forward(x)
x = self.norm2(x + self.dropout(ffn_out))
return x解码器层
python
class TransformerDecoderLayer(nn.Module):
"""Transformer 解码器层
结构:
1. 带因果掩码的自注意力
2. 交叉注意力 (关注 Encoder 输出)
3. FFN
每步后都有残差+Norm
"""
def __init__(self, d_model: int = 512, n_heads: int = 8,
d_ff: int = 2048, dropout: float = 0.1,
activation: str = 'relu'):
super().__init__()
# 1. 掩码自注意力(因果掩码,只看已经生成的内容)
self.masked_self_attn = MultiHeadAttention(d_model, n_heads, dropout)
# 2. 交叉注意力(Q 来自 Decoder, K,V 来自 Encoder)
self.cross_attn = MultiHeadCrossAttention(d_model, n_heads, dropout)
# 3. FFN
self.feed_forward = FeedForward(d_model, d_ff, dropout, activation)
# 三个 LayerNorm
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.norm3 = nn.LayerNorm(d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x: torch.Tensor, encoder_output: torch.Tensor,
tgt_mask: torch.Tensor = None,
memory_mask: torch.Tensor = None) -> torch.Tensor:
"""
Args:
x: (batch, tgt_len, d_model) — Decoder 输入
encoder_output: (batch, src_len, d_model) — Encoder 输出
tgt_mask: 因果掩码 (tgt_len, tgt_len)
memory_mask: Encoder padding mask (batch, 1, 1, src_len)
"""
# 子层1: 掩码自注意力
attn_out = self.masked_self_attn(x, x, x, tgt_mask)
x = self.norm1(x + self.dropout(attn_out))
# 子层2: 交叉注意力
cross_out = self.cross_attn(x, encoder_output, memory_mask)
x = self.norm2(x + self.dropout(cross_out))
# 子层3: FFN
ffn_out = self.feed_forward(x)
x = self.norm3(x + self.dropout(ffn_out))
return x8. 完整 Transformer 实现
python
class Transformer(nn.Module):
"""
完整的 Transformer 模型 (Encoder-Decoder)
论文: Attention Is All You Need (Vaswani et al., 2017)
架构参数:
- d_model = 512 (模型维度)
- n_heads = 8 (注意力头数)
- d_ff = 2048 (FFN 中间维度)
- n_layers = 6 (编码器/解码器层数)
- dropout = 0.1 (正则化)
"""
def __init__(self, src_vocab_size: int, tgt_vocab_size: int,
d_model: int = 512, n_heads: int = 8,
d_ff: int = 2048, n_layers: int = 6,
dropout: float = 0.1, max_len: int = 5000):
super().__init__()
self.d_model = d_model # 保存下来,供 encode/decode 方法使用
self.n_heads = n_heads
# 编码器组件
self.src_embedding = nn.Embedding(src_vocab_size, d_model)
self.src_positional = SinusoidalPositionalEncoding(d_model, max_len, dropout)
self.encoder_layers = nn.ModuleList([
TransformerEncoderLayer(d_model, n_heads, d_ff, dropout)
for _ in range(n_layers)
])
# 解码器组件
self.tgt_embedding = nn.Embedding(tgt_vocab_size, d_model)
self.tgt_positional = SinusoidalPositionalEncoding(d_model, max_len, dropout)
self.decoder_layers = nn.ModuleList([
TransformerDecoderLayer(d_model, n_heads, d_ff, dropout)
for _ in range(n_layers)
])
# 输出投影: d_model → vocab_size
self.output_projection = nn.Linear(d_model, tgt_vocab_size)
# 参数初始化
self._init_parameters()
def _init_parameters(self):
"""Xavier/Glorot 初始化"""
for p in self.parameters():
if p.dim() > 1:
nn.init.xavier_uniform_(p)
def encode(self, src: torch.Tensor,
src_mask: torch.Tensor = None) -> torch.Tensor:
"""编码器"""
x = self.src_embedding(src) * math.sqrt(self.d_model)
x = self.src_positional(x)
for layer in self.encoder_layers:
x = layer(x, src_mask)
return x
def decode(self, tgt: torch.Tensor, memory: torch.Tensor,
tgt_mask: torch.Tensor = None,
memory_mask: torch.Tensor = None) -> torch.Tensor:
"""解码器"""
x = self.tgt_embedding(tgt) * math.sqrt(self.d_model)
x = self.tgt_positional(x)
for layer in self.decoder_layers:
x = layer(x, memory, tgt_mask, memory_mask)
return x
def forward(self, src: torch.Tensor, tgt: torch.Tensor,
src_mask: torch.Tensor = None,
tgt_mask: torch.Tensor = None) -> torch.Tensor:
"""
Args:
src: (batch, src_len) — 源序列 token IDs
tgt: (batch, tgt_len) — 目标序列 token IDs
src_mask: padding mask
tgt_mask: 因果掩码 + padding mask
Returns:
logits: (batch, tgt_len, tgt_vocab_size)
"""
# 编码
memory = self.encode(src, src_mask)
# 解码
output = self.decode(tgt, memory, tgt_mask, src_mask)
# 投影到词表大小
return self.output_projection(output)9. Transformer 训练与推理
9.1 损失函数
Transformer 使用标签平滑交叉熵:
标签平滑 (
9.2 优化器与学习率调度
原论文使用 Adam 优化器配合自定义的 warmup 调度:
python
class TransformerLRScheduler:
"""原论文的学习率调度器"""
def __init__(self, optimizer, d_model: int, warmup_steps: int = 4000):
self.optimizer = optimizer
self.d_model = d_model
self.warmup_steps = warmup_steps
self.step_num = 0
def step(self):
self.step_num += 1
lr = self.d_model ** (-0.5) * min(
self.step_num ** (-0.5),
self.step_num * self.warmup_steps ** (-1.5)
)
for param_group in self.optimizer.param_groups:
param_group['lr'] = lr
return lr10. 计算复杂度分析
设序列长度为
| 操作 | 时间复杂度 | 内存复杂度 | 说明 |
|---|---|---|---|
| 自注意力 | 注意力矩阵 | ||
| FFN | 逐位置的全连接 | ||
| 总体 | 当 |
这也是为什么长上下文(
很大)如此昂贵。Flash Attention、Sparse Attention、Linear Attention 等方法都在试图降低 复杂度。
Flash Attention
Flash Attention 通过分块计算和 IO 优化,在不改变数学结果的前提下,将内存复杂度从
python
# 使用 Flash Attention (PyTorch 2.0+)
# 只需将 attention 替换为:
# F.scaled_dot_product_attention(Q, K, V, is_causal=True)
# PyTorch 会自动使用 Flash Attention 如果可用11. Transformer 的现代变体
mermaid
graph TD
ROOT["原版 Transformer<br/>Vaswani et al., 2017"] --> BERT["BERT<br/>Encoder-Only<br/>双向、MLM预训练"]
ROOT --> GPT["GPT<br/>Decoder-Only<br/>单向、CLM预训练"]
ROOT --> T5["T5<br/>Encoder-Decoder<br/>Span Corruption"]
GPT --> GPT3["GPT-3 (2020)<br/>175B, In-Context Learning"]
GPT3 --> LLAMA["LLaMA (2023)<br/>RoPE, SwiGLU, RMSNorm"]
LLAMA --> LLAMA2["LLaMA 2 (2023)<br/>GQA, 2T tokens"]
LLAMA2 --> LLAMA3["LLaMA 3 (2024)<br/>128K vocab, 15T tokens"]
LLAMA3 --> LLAMA4["LLaMA 4 (2025)<br/>MoE 架构, 原生多模态"]
LLAMA3 --> DS["DeepSeek-V3/R2 (2025)<br/>MoE + MLA + 推理增强"]
LLAMA3 --> QWEN["Qwen4 (2025)<br/>MoE + Thinking 模式"]
BERT --> ROBERTA["RoBERTa<br/>更好的训练配方"]
BERT --> DEBERTA["DeBERTa<br/>解耦注意力"]
style LLAMA4 fill:#e74c3c,color:#fff
style DS fill:#9b59b6,color:#fff
style GPT3 fill:#f39c12,color:#fff现代 Transformer 的标配改进
| 改进 | 说明 | 使用模型 |
|---|---|---|
| Pre-LN | 先 Norm 再子层 | GPT-2+, LLaMA |
| RMSNorm | 去掉均值的简化 LayerNorm | LLaMA, Mistral |
| RoPE | 旋转位置编码 | LLaMA, Qwen, Mistral |
| SwiGLU | 门控 FFN | LLaMA, PaLM |
| GQA/MQA | 分组/多查询注意力 | LLaMA 2+, Mistral |
| Flash Attention | IO 优化的注意力 | 几乎所有现代框架 |
| Untied Embedding | 输入/输出不共享 Embedding | GPT-3, LLaMA |
| MoE | 混合专家,稀疏激活部分参数 | LLaMA 4, DeepSeek-V3, Qwen4 |
| MLA | 多头潜在注意力,KV Cache 压缩 | DeepSeek-V3/R2 |
12. MoE (Mixture of Experts) — 稀疏激活的万亿参数架构
为什么需要 MoE?
Dense 模型的困境:参数量翻倍 → 计算量翻倍 → 推理成本翻倍。MoE 打破了这个等式:参数量翻倍,但每次推理只激活一小部分参数。
mermaid
graph TD
subgraph "Dense Model (70B)"
IN1["输入 Token"] --> ALL["全部 70B 参数<br/>每次都参与计算"]
ALL --> OUT1["输出"]
end
subgraph "MoE Model (400B 总参数, 激活 ~50B)"
IN2["输入 Token"] --> ROUTER["Router<br/>门控网络"]
ROUTER -->|"选择 Top-2"| E1["Expert 1<br/>(50B)"]
ROUTER -->|"选择 Top-2"| E3["Expert 3<br/>(50B)"]
ROUTER -.->|"未选中"| E2["Expert 2<br/>(50B)"]
ROUTER -.->|"未选中"| E4["Expert 4<br/>(50B)"]
ROUTER -.->|"未选中"| E5["Expert 5<br/>(50B)"]
ROUTER -.->|"未选中"| E6["Expert 6<br/>(50B)"]
ROUTER -.->|"未选中"| E7["Expert 7<br/>(50B)"]
ROUTER -.->|"未选中"| E8["Expert 8<br/>(50B)"]
E1 --> MIX["加权混合"]
E3 --> MIX
MIX --> OUT2["输出"]
end
style ROUTER fill:#e74c3c,color:#fff
style E1 fill:#2ecc71,color:#fff
style E3 fill:#2ecc71,color:#fffMoE 核心组件实现
python
import torch
import torch.nn as nn
import torch.nn.functional as F
class Expert(nn.Module):
"""单个专家 — 本质上就是一个 FFN"""
def __init__(self, d_model: int, d_ff: int):
super().__init__()
self.w1 = nn.Linear(d_model, d_ff, bias=False)
self.w2 = nn.Linear(d_ff, d_model, bias=False)
self.w3 = nn.Linear(d_model, d_ff, bias=False) # SwiGLU gate
def forward(self, x):
# SwiGLU: gate * up
return self.w2(F.silu(self.w1(x)) * self.w3(x))
class TopKRouter(nn.Module):
"""Top-K 路由器 — 决定每个 Token 发给哪些专家"""
def __init__(self, d_model: int, num_experts: int, top_k: int = 2):
super().__init__()
self.top_k = top_k
self.gate = nn.Linear(d_model, num_experts, bias=False)
def forward(self, x):
# x: [batch, seq_len, d_model]
logits = self.gate(x) # [batch, seq, num_experts]
scores = F.softmax(logits, dim=-1)
# 选择 Top-K 专家
top_k_scores, top_k_indices = torch.topk(scores, self.top_k, dim=-1)
# 归一化权重(使 Top-K 权重之和为 1)
top_k_scores = top_k_scores / top_k_scores.sum(dim=-1, keepdim=True)
return top_k_scores, top_k_indices
class MoELayer(nn.Module):
"""完整的 MoE 层 — 替代 Transformer 中的 FFN"""
def __init__(
self,
d_model: int = 4096,
d_ff: int = 11008,
num_experts: int = 8,
top_k: int = 2,
num_shared_experts: int = 1, # DeepSeek 风格的共享专家
):
super().__init__()
self.num_experts = num_experts
self.top_k = top_k
self.num_shared_experts = num_shared_experts
# 路由器
self.router = TopKRouter(d_model, num_experts, top_k)
# 专家池
self.experts = nn.ModuleList([
Expert(d_model, d_ff) for _ in range(num_experts)
])
# 共享专家(所有 Token 都经过,保证基础能力)
if num_shared_experts > 0:
self.shared_experts = nn.ModuleList([
Expert(d_model, d_ff) for _ in range(num_shared_experts)
])
def forward(self, x):
batch_size, seq_len, d_model = x.shape
# 1. 路由
scores, indices = self.router(x) # [B, S, top_k], [B, S, top_k]
# 2. 分发到各专家并加权求和
output = torch.zeros_like(x)
flat_x = x.view(-1, d_model) # [B*S, D]
flat_scores = scores.view(-1, self.top_k) # [B*S, top_k]
flat_indices = indices.view(-1, self.top_k) # [B*S, top_k]
for k in range(self.top_k):
expert_idx = flat_indices[:, k] # [B*S]
weight = flat_scores[:, k].unsqueeze(-1) # [B*S, 1]
for i in range(self.num_experts):
mask = (expert_idx == i)
if mask.any():
expert_input = flat_x[mask]
expert_output = self.experts[i](expert_input)
output.view(-1, d_model)[mask] += weight[mask] * expert_output
# 3. 共享专家(所有 Token 都经过)
if self.num_shared_experts > 0:
for shared_expert in self.shared_experts:
output = output + shared_expert(x)
return output
# 使用示例
moe = MoELayer(d_model=4096, d_ff=11008, num_experts=8, top_k=2)
x = torch.randn(2, 128, 4096)
out = moe(x)
print(f"输入: {x.shape} → 输出: {out.shape}")
print(f"总参数: {sum(p.numel() for p in moe.parameters()) / 1e6:.1f}M")
print(f"每 Token 激活参数: ~{sum(p.numel() for p in moe.parameters()) / 1e6 * 2/8:.1f}M (Top-2/8)")MoE 的工程挑战
| 挑战 | 问题 | 解决方案 |
|---|---|---|
| 负载不均衡 | 某些专家被过度选择,其他专家闲置 | 辅助损失(Auxiliary Loss)惩罚不均衡 |
| 通信开销 | 多卡部署时 Token 需要跨 GPU 发送到对应专家 | Expert Parallelism + All-to-All 通信 |
| 显存占用 | 虽然计算量少,但所有专家参数都要加载 | 量化 + Offloading |
| 训练不稳定 | 路由器梯度稀疏,容易坍塌 | Z-Loss 正则化 + Jitter Noise |
MoE 代表模型对比
| 模型 | 总参数 | 激活参数 | 专家数 | Top-K | 特色 |
|---|---|---|---|---|---|
| LLaMA 4 Scout | 109B | 17B | 16 | 2 | 原生多模态 MoE |
| DeepSeek-V3 | 671B | 37B | 256 | 8 | MLA + 共享专家 |
| Qwen4-MoE | 57B | 14B | 64 | 8 | 细粒度专家 |
| Mixtral 8x7B | 47B | 13B | 8 | 2 | 最早的开源 MoE |
13. MLA (Multi-head Latent Attention) — KV Cache 压缩
DeepSeek-V3/R2 的核心创新:将 KV Cache 压缩到低维潜在空间,大幅减少推理时的显存占用。
python
class MultiHeadLatentAttention(nn.Module):
"""
MLA: 将 K/V 投影到低维空间再存储
标准 MHA: KV Cache = 2 × n_heads × d_head × seq_len
MLA: KV Cache = d_compress × seq_len (压缩 4-8x)
"""
def __init__(
self,
d_model: int = 4096,
n_heads: int = 32,
d_head: int = 128,
d_compress: int = 512, # 压缩维度(远小于 n_heads × d_head)
):
super().__init__()
self.n_heads = n_heads
self.d_head = d_head
self.d_compress = d_compress
# 下投影:将 KV 压缩到低维
self.kv_down_proj = nn.Linear(d_model, d_compress, bias=False)
# 上投影:从低维恢复 K 和 V
self.k_up_proj = nn.Linear(d_compress, n_heads * d_head, bias=False)
self.v_up_proj = nn.Linear(d_compress, n_heads * d_head, bias=False)
# Q 正常投影
self.q_proj = nn.Linear(d_model, n_heads * d_head, bias=False)
self.o_proj = nn.Linear(n_heads * d_head, d_model, bias=False)
self.scale = d_head ** -0.5
def forward(self, x, kv_cache=None):
B, S, _ = x.shape
# Q 正常计算
Q = self.q_proj(x).view(B, S, self.n_heads, self.d_head).transpose(1, 2)
# KV 先压缩再恢复
compressed_kv = self.kv_down_proj(x) # [B, S, d_compress] ← 只缓存这个!
K = self.k_up_proj(compressed_kv).view(B, S, self.n_heads, self.d_head).transpose(1, 2)
V = self.v_up_proj(compressed_kv).view(B, S, self.n_heads, self.d_head).transpose(1, 2)
# 标准注意力计算
attn = (Q @ K.transpose(-2, -1)) * self.scale
attn = F.softmax(attn, dim=-1)
out = attn @ V
out = out.transpose(1, 2).contiguous().view(B, S, -1)
return self.o_proj(out)
# KV Cache 对比
d_model, n_heads, d_head, seq_len = 4096, 32, 128, 4096
standard_kv_cache = 2 * n_heads * d_head * seq_len * 2 # FP16, bytes
mla_kv_cache = 512 * seq_len * 2 # d_compress=512, FP16
print(f"标准 MHA KV Cache: {standard_kv_cache / 1024 / 1024:.1f} MB/layer")
print(f"MLA KV Cache: {mla_kv_cache / 1024 / 1024:.1f} MB/layer")
print(f"压缩比: {standard_kv_cache / mla_kv_cache:.1f}x")参考
- Attention Is All You Need — Vaswani et al., 2017
- The Annotated Transformer — Harvard NLP 的逐行注释实现
- The Illustrated Transformer — Jay Alammar 的可视化讲解
- LLaMA: Open and Efficient Foundation Language Models — Touvron et al., 2023
- The LLaMA 4 Herd of Models — Meta AI, 2025
- DeepSeek-V3 Technical Report — DeepSeek-AI, 2024
- DeepSeek-R2: Incentivizing Reasoning Capability in LLMs via RL — DeepSeek-AI, 2025
登录后即可发表评论 👇