Skip to content

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 模型的工作方式是:

  1. Encoder RNN 把整个输入压缩成一个固定长度的向量
  2. 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:#fff

GPT 只用了 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 映射为稠密向量:

E∈RV×dmodel

其中 V 是词汇量(通常 30K-50K),dmodel 是模型维度(原论文 dmodel=512)。

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" 对模型来说是一样的。

原论文使用正弦/余弦位置编码:

PE(pos,2i)=sin⁡(pos100002i/dmodel)PE(pos,2i+1)=cos⁡(pos100002i/dmodel)

其中:

  • pos:词在序列中的位置(0, 1, 2, ...)
  • i:维度索引(0, 1, ..., d_model/2 - 1)
  • dmodel:模型维度

关键性质:

  1. 每个位置的编码是唯一的
  2. 相邻位置的编码相似(位置 5 和 6 的编码比 5 和 50 更相似)
  3. 不同维度对应不同频率——低维度编码"粗粒度"位置,高维度编码"细粒度"位置
  4. PEpos+k 可以由 PEpos 的线性函数表示 → 模型能学到相对位置关系
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可学习的 EmbeddingBERT, GPT, ViT
RoPE (旋转位置编码)通过旋转矩阵编码相对位置LLaMA, Qwen, Mistral
ALiBi在 Attention 分数上加线性偏置BLOOM
相对位置编码编码词对之间的相对距离Transformer-XL

RoPE 是当前 LLM 的首选。它将位置信息编码为旋转操作,天然处理相对位置,支持任意长度的外推。

RoPE 的核心思想:

qm⊤kn=(RΘ,mq)⊤(RΘ,nk)=q⊤RΘ,n−mk

即旋转后的 Q 和 K 的点积只依赖于它们的相对位置 n−m。


3. 自注意力机制 — Transformer 的核心 ​

3.1 从直觉出发:什么是"注意力"? ​

在正式公式之前,先用一个例子理解什么是"查询-键-值"(Query-Key-Value):

想象你正在阅读下面这句话:

"猫追老鼠因为它饿了"

你的大脑需要判断"它"指代什么。大脑怎么做?

  1. 生成一个 查询(Query):"谁有可能被'它'指代?" → 查询包含代词的特征(第三人称、单数、动物)
  2. 句子中每个词都有一个 键(Key):描述"我是什么类型的词"
    • "猫"的 Key:[动物, 主语, 第三人称]
    • "追"的 Key:[动作, 谓语, ...]
    • "老鼠"的 Key:[动物, 宾语, 第三人称]
  3. 查询与键匹配:Q("它") 与 K("猫") 的点积 = 相似度分数。谁的分数高就说明 Q 更"关注"谁
  4. 每个词还有一个 值(Value):表示"如果我被选中,我要提供什么信息"
  5. 加权求和:用匹配分数做权重,对 Value 加权平均 → 得到最终输出

公式表达的就是这个"查字典"的过程:

Attention(Q,K,V)=softmax(QKTdk)⋅V

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

公式:

Attention(Q,K,V)=softmax(QKTdk)⋅V
步骤操作形状变化直觉
1QKT(seq, d_k) × (d_k, seq) → (seq, seq)每个词和每个词"聊天",看谁和谁相关
2/dk(seq, seq)控制数值大小,别让 softmax 饱和
3softmax(seq, seq)把分数变成"谁占多少权重"
4×V(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_weights

3.3 因果掩码 (Causal Mask) — GPT 的关键 ​

GPT 是自回归模型,生成第 t 个 Token 时只能看到前面 t−1 个 Token。因果掩码实现:

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 为什么需要多头? ​

单头注意力的表达能力有限。多头意味着:

  • 不同的"注意力头"关注不同类型的依赖关系
  • 有的头关注语法关系(主语-谓语),有的关注语义关系(代词-指代对象)
  • 每个头的维度更小(dk=dmodel/h),总计算量不变
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

公式:

MultiHead(Q,K,V)=Concat(head1,...,headh)⋅WOwhere headi=Attention(QWiQ,KWiK,VWiV)

其中 WiQ,WiK∈Rdmodel×dk,WiV∈Rdmodel×dv,WO∈Rhdv×dmodel。

原论文取 h=8, dk=dv=dmodel/h=64。

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 节省质量损失使用模型
MHAhh基准无原版 Transformer, BERT
GQAhg (< h)h/g 倍几乎无LLaMA 2/3 (g=8), Mistral
MQAh1h 倍轻微PaLM, Gemini

5. 前馈网络 (Feed-Forward Network) ​

5.1 标准 FFN ​

每个位置独立应用的两层全连接网络:

FFN(x)=max(0,xW1+b1)W2+b2
  • 输入维度:dmodel (512)
  • 中间维度:dff (2048, 4×)
  • 输出维度:dmodel (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 等模型的标配:

SwiGLU(x)=(xW1⊙SiLU(xWg))W2

其中 SiLU(x)=x⋅σ(x)。门控机制让模型学会"选择性忽略"不重要的信息。


6. 残差连接与层归一化 ​

6.1 残差连接 (Residual Connection) ​

Output=LayerNorm(x+Sublayer(x))

残差连接让梯度可以"直接流过"每一层,是训练深层网络的关键:

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 按特征维度归一化:

LayerNorm(x)=γ⋅x−μσ2+ϵ+β

其中:

  • μ=1dmodel∑i=1dmodelxi(均值)
  • σ2=1dmodel∑i=1dmodel(xi−μ)2(方差)
  • γ,β 是可学习的缩放和偏移参数
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?

  1. BatchNorm 依赖 batch 统计量,batch 太小不稳定
  2. 序列长度可变时,BatchNorm 需要处理 padding
  3. LayerNorm 在推理时不依赖 batch,行为完全确定

Pre-LN vs Post-LN ​

现代 LLM 普遍使用 Pre-LN(先 Norm 再子层),训练更稳定:

变体公式特点
Post-LN (原版)LN(x+Sublayer(x))原论文实现,需要 warmup
Pre-LN (现代)x+Sublayer(LN(x))训练更稳定,无需 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 x

8. 完整 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 使用标签平滑交叉熵:

L=−∑t=1T∑v=1V[(1−ϵ)⋅1yt=v+ϵV]⋅log⁡P(v|y<t)

标签平滑 (ϵ=0.1) 防止模型过于自信,提高泛化能力。

9.2 优化器与学习率调度 ​

原论文使用 Adam 优化器配合自定义的 warmup 调度:

lr=dmodel−0.5⋅min(step_num−0.5,step_num⋅warmup_steps−1.5)
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 lr

10. 计算复杂度分析 ​

设序列长度为 n, 模型维度为 d:

操作时间复杂度内存复杂度说明
自注意力O(n2⋅d)O(n2)注意力矩阵 n×n
FFNO(n⋅d2)O(n⋅d)逐位置的全连接
总体O(n2⋅d+n⋅d2)O(n2)当 n≫d 时注意力是瓶颈

这也是为什么长上下文(n 很大)如此昂贵。Flash Attention、Sparse Attention、Linear Attention 等方法都在试图降低 O(n2) 复杂度。

Flash Attention ​

Flash Attention 通过分块计算和 IO 优化,在不改变数学结果的前提下,将内存复杂度从 O(n2) 降到 O(n)(以 blocks 为单位),速度提升 2-4x:

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去掉均值的简化 LayerNormLLaMA, Mistral
RoPE旋转位置编码LLaMA, Qwen, Mistral
SwiGLU门控 FFNLLaMA, PaLM
GQA/MQA分组/多查询注意力LLaMA 2+, Mistral
Flash AttentionIO 优化的注意力几乎所有现代框架
Untied Embedding输入/输出不共享 EmbeddingGPT-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:#fff

MoE 核心组件实现 ​

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 Scout109B17B162原生多模态 MoE
DeepSeek-V3671B37B2568MLA + 共享专家
Qwen4-MoE57B14B648细粒度专家
Mixtral 8x7B47B13B82最早的开源 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")

参考 ​

批注模式

💬 文章评论

暂无评论,来说点什么吧 👇

编程学习笔记