从零实现Transformer:手写自注意力到GPT-2级别语言模型
一、引言
2017年,"Attention Is All You Need" 论文以一句"Attention Is All You Need"宣告了 NLP 新时代的到来。Transformer 不依赖 RNN 的序列递归,用纯粹的注意力机制实现了并行训练,支撑了 GPT-4、Claude、Gemini 等所有现代大模型。
本文将带你从零实现一个完整的 Transformer :从自注意力的数学推导开始,手写每一行代码,最终训练一个 GPT-2 级别的语言模型。我们会逐一拆解:Scaled Dot-Product Attention 的 QK^T/√d_k 为什么除以 √d_k?Multi-Head Attention 为什么有效?Positional Encoding 的 sin/cos 设计原理?Layer Normalization 为什么放在前面(Pre-Norm)?最后的 RoPE 和 FlashAttention 又是如何进化的?
这是一篇"看完能自己写出 Transformer"的深度文章。
二、自注意力:从直觉到数学
2.1 为什么需要注意力
考虑句子:"The cat sat on the mat because it was tired."
"it" 指代 "cat" 还是 "mat"?人类通过上下文语义关联来判断------"tired"(累)更可能描述有生命的 "cat" 而不是 "mat"。RNN 处理这个问题时,需要将"cat"的信息通过隐藏状态传递6个时间步到"it"位置,信息在传递中衰减。Transformer 的解决方案是:让每个词直接看到所有其他词,并通过可学习的相似度来决定"关注"多少。
2.2 数学定义
自注意力的核心公式:
Attention(Q, K, V) = softmax(QK^T / √d_k) V
其中 Q(Query)、K(Key)、V(Value)都来自同一个输入 X 的线性变换:
Q = XW_Q, K = XW_K, V = XW_V
直觉理解:
- Query(查询):"我(当前词)想找什么?"
- Key(键):"我(其他词)能提供什么?"
- Value(值):"我(其他词)的实际内容是什么?"
- Q·K^T:查询和键的相似度 → 注意力权重
- √d_k:缩放因子,防止点积过大导致 softmax 梯度消失
- softmax:归一化为概率分布
- × V:用注意力权重加权聚合 Value
2.3 为什么除以 √d_k?
这是面试高频题。假设 Q 和 K 的每个元素是均值为0、方差为1的独立随机变量,那么点积 q·k = Σq_i·k_i 的方差是 d_k(每项方差为1,共 d_k 项独立相加)。
当 d_k=64 时,q·k 的标准差为 8。这意味着:
- 大约有 1/3 的点积值落在 -8, 8 之外
- 经过 softmax 后,这些极端值会形成接近 one-hot 的分布
- 梯度接近零 → 训练停滞
除以 √d_k 将方差归一化到 1,使 softmax 输入保持在合理范围,梯度流畅。
2.4 手写实现
python
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
class ScaledDotProductAttention(nn.Module):
"""单头缩放点积注意力------Transformer的原子操作"""
def __init__(self, d_k=64):
super().__init__()
self.d_k = d_k
def forward(self, Q, K, V, mask=None):
"""
Q, K, V: [batch_size, seq_len, d_k]
mask: [batch_size, seq_len, seq_len] 或 [batch_size, 1, 1, seq_len]
返回: output [batch_size, seq_len, d_k], attention_weights
"""
# 1. 计算注意力分数: scores = Q @ K^T / √d_k
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
# scores: [batch_size, seq_len, seq_len]
# scores[i, j, k] = 第i个样本中,位置j对位置k的注意力分数
# 2. Mask处理(因果mask或padding mask)
if mask is not None:
# mask中True的位置设为 -inf(经过softmax后变为0)
scores = scores.masked_fill(mask == 0, float('-inf'))
# 3. Softmax归一化
attention_weights = F.softmax(scores, dim=-1)
# 每一行的和为1: 当前位置对所有位置的注意力权重
# 4. Dropout正则化(训练时)
attention_weights = F.dropout(attention_weights, p=0.1, training=self.training)
# 5. 加权聚合Value
output = torch.matmul(attention_weights, V)
return output, attention_weights
# ============ 可视化注意力矩阵 ============
# 构造一个简单例子
d_k = 64
seq_len = 8
batch_size = 2
# 随机输入(模拟词嵌入)
Q = torch.randn(batch_size, seq_len, d_k)
K = torch.randn(batch_size, seq_len, d_k)
V = torch.randn(batch_size, seq_len, d_k)
# 创建一个因果mask(左下三角为1,右上为0)
# 确保位置i只能看到位置0到i
causal_mask = torch.tril(torch.ones(seq_len, seq_len)).unsqueeze(0).unsqueeze(0)
# causal_mask: [1, 1, seq_len, seq_len]
attention = ScaledDotProductAttention(d_k=d_k)
output, weights = attention(Q, K, V, mask=causal_mask)
print(f"Input shape: {Q.shape}")
print(f"Output shape: {output.shape}")
print(f"Attention weights shape: {weights.shape}")
print(f"\nAttention pattern (sample 0, row sums should be 1):")
print(f" {weights[0, 0].sum(dim=-1).tolist()}") # 验证每行和为1
print(f"\nCausal mask pattern (1=visible, 0=masked):")
for i in range(8):
row = ['1' if x > 0.5 else '0' for x in weights[0, 0, i].tolist()]
print(f" pos{i}: {''.join(row)}")
2.5 可视化:因果注意力模式
运行上面的代码,你会看到清晰的下三角注意力模式:
pos0: 10000000 ← token 0 只能看到自己
pos1: 11000000 ← token 1 可以看到0和1
pos2: 11100000
pos3: 11110000
pos4: 11111000
pos5: 11111100
pos6: 11111110
pos7: 11111111 ← token 7 可以看到所有之前的token
这正是 GPT 的自回归生成模式------每个token只能关注它自己和之前的位置,不能"偷看"未来。
三、多头注意力:为什么需要多个头
3.1 直觉
单头注意力只有一个"视角"。对于句子 "I love you":
- 头1 可能关注句法关系(主谓宾)
- 头2 可能关注语义关系(love→情感词)
- 头3 可能关注位置关系(相邻词)
- 头4 可能关注指代关系(代词→先行词)
多个头并行计算,每个头关注不同的特征子空间,然后拼接起来。
3.2 手写实现
python
class MultiHeadAttention(nn.Module):
"""多头注意力------并行计算多组QKV"""
def __init__(self, d_model=512, n_heads=8, dropout=0.1):
super().__init__()
assert d_model % n_heads == 0, "d_model必须能被n_heads整除"
self.d_model = d_model
self.n_heads = n_heads
self.d_k = d_model // n_heads # 每个头的维度
# Q、K、V的联合投影(一次矩阵乘法代替三次)
self.W_qkv = nn.Linear(d_model, 3 * d_model, bias=False)
# 输出投影
self.W_o = nn.Linear(d_model, d_model)
self.attention = ScaledDotProductAttention(d_k=self.d_k)
self.dropout = nn.Dropout(dropout)
def split_heads(self, x):
"""
将 [batch_size, seq_len, d_model]
拆分为 [batch_size, n_heads, seq_len, d_k]
这样每个头独立计算注意力
"""
batch_size, seq_len, _ = x.shape
x = x.view(batch_size, seq_len, self.n_heads, self.d_k)
return x.transpose(1, 2) # [batch_size, n_heads, seq_len, d_k]
def combine_heads(self, x):
"""逆操作:合并多头 → [batch_size, seq_len, d_model]"""
batch_size, _, seq_len, _ = x.shape
x = x.transpose(1, 2).contiguous()
return x.view(batch_size, seq_len, self.d_model)
def forward(self, x, mask=None):
batch_size, seq_len, _ = x.shape
# 1. 联合投影: [batch, seq, d_model] → [batch, seq, 3*d_model]
qkv = self.W_qkv(x)
# 2. 拆分为 Q, K, V
Q, K, V = qkv.chunk(3, dim=-1)
# 3. 拆分为多头
Q = self.split_heads(Q) # [batch, n_heads, seq, d_k]
K = self.split_heads(K)
V = self.split_heads(V)
# 4. 并行计算注意力(每个头独立)
# mask需要扩展head维度: [batch, 1, 1, seq] → [batch, 1, seq, seq]
if mask is not None:
mask = mask.unsqueeze(1) # 添加head维度
# 5. 缩放点积注意力
attn_output, attn_weights = self.attention(Q, K, V, mask)
# 6. 合并多头
output = self.combine_heads(attn_output)
# 7. 输出投影
output = self.W_o(output)
return output, attn_weights
3.3 参数量分析
假设 d_model=512, n_heads=8:
- W_qkv: 512 × (3×512) = 786,432 参数
- W_o: 512 × 512 = 262,144 参数
- 总计约 1M 参数
加上后续的 FFN,一个 Transformer Block 约 7M 参数。GPT-2 Small(12层)约 124M 参数,与手算一致。
四、位置编码:sin/cos的优雅设计
4.1 为什么需要位置编码
注意力机制是置换等变的------交换输入序列中的两个位置,输出也会相应交换(但不改变值)。这意味着"ABCD"和"DCBA"的注意力输出只是顺序不同,内容完全相同。但语言是有顺序的------"狗咬人"和"人咬狗"语义完全不同。
位置编码(Positional Encoding)为每个位置注入位置信息,打破置换等变性。
4.2 正弦位置编码
python
class SinusoidalPositionalEncoding(nn.Module):
"""原始Transformer的sin/cos位置编码------无需学习参数"""
def __init__(self, d_model=512, max_len=5000):
super().__init__()
# 创建位置编码矩阵 [max_len, d_model]
pe = torch.zeros(max_len, d_model)
# 位置索引 [max_len, 1]
position = torch.arange(0, max_len).unsqueeze(1).float()
# 频率项: 维度越高的特征频率越低(波长从2π到10000·2π)
div_term = torch.exp(
torch.arange(0, d_model, 2).float() *
(-math.log(10000.0) / d_model)
)
# 偶数列用sin,奇数列用cos
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
# 添加batch维度
pe = pe.unsqueeze(0) # [1, max_len, d_model]
# 注册为buffer(不参与梯度,但随模型保存/加载)
self.register_buffer('pe', pe)
def forward(self, x):
"""x: [batch_size, seq_len, d_model]"""
return x + self.pe[:, :x.size(1), :]
# 可视化不同位置的位置编码相似度
import matplotlib.pyplot as plt
import numpy as np
def visualize_position_similarity(d_model=128, max_len=100):
pe = SinusoidalPositionalEncoding(d_model, max_len)
pe_matrix = pe.pe[0].numpy() # [max_len, d_model]
# 计算位置间的余弦相似度
similarity = np.zeros((max_len, max_len))
for i in range(max_len):
for j in range(max_len):
similarity[i, j] = np.dot(pe_matrix[i], pe_matrix[j]) / (
np.linalg.norm(pe_matrix[i]) * np.linalg.norm(pe_matrix[j])
)
plt.figure(figsize=(10, 8))
plt.imshow(similarity, cmap='RdBu_r', aspect='auto')
plt.colorbar(label='Cosine Similarity')
plt.xlabel('Position j'); plt.ylabel('Position i')
plt.title('Positional Encoding: Cosine Similarity Between Positions')
plt.show()
# 关键观察:
# 1. 对角线=1(自己和自己最相似)
# 2. 相邻位置相似度高(浅色带)
# 3. 远离位置相似度低甚至负(可以区分远近)
# 关键特性:
# - PE(pos+k) 可以表示为 PE(pos) 的线性函数
# (sin(a+b)=sin(a)cos(b)+cos(a)sin(b))
# 这使得模型可以学习"相对位置"而非"绝对位置"
# - 无需训练参数,支持外推到训练时未见过的序列长度
五、Feed-Forward Network
python
class FeedForward(nn.Module):
"""Position-wise FFN------每个token独立通过同样的2层MLP"""
def __init__(self, d_model=512, d_ff=2048, dropout=0.1):
super().__init__()
self.linear1 = nn.Linear(d_model, d_ff) # 扩展4倍
self.linear2 = nn.Linear(d_ff, d_model) # 压缩回原维度
self.dropout = nn.Dropout(dropout)
def forward(self, x):
# GELU激活(GPT使用)vs ReLU(原论文使用)
return self.linear2(self.dropout(F.gelu(self.linear1(x))))
# 为什么扩展4倍?
# FFN提供非线性变换能力,扩展维度增大表示容量
# 类似于CNN中的1×1卷积 → 增加通道数
六、Layer Normalization
python
class TransformerBlock(nn.Module):
"""一个完整的Transformer层(GPT-2风格: Pre-Norm)"""
def __init__(self, d_model=512, n_heads=8, d_ff=2048, dropout=0.1):
super().__init__()
# Pre-Norm: LayerNorm在子层之前(现代标准做法)
self.ln1 = nn.LayerNorm(d_model)
self.ln2 = nn.LayerNorm(d_model)
self.attention = MultiHeadAttention(d_model, n_heads, dropout)
self.ffn = FeedForward(d_model, d_ff, dropout)
self.dropout = nn.Dropout(dropout)
def forward(self, x, mask=None):
# 1. 多头注意力 + 残差连接
attn_out, attn_weights = self.attention(self.ln1(x), mask)
x = x + self.dropout(attn_out) # ★ 残差连接
# 2. FFN + 残差连接
ffn_out = self.ffn(self.ln2(x))
x = x + self.dropout(ffn_out)
return x, attn_weights
# Pre-Norm vs Post-Norm (历史演进):
# Post-Norm(原始): x + Sublayer(LN(x)) ← 梯度可能爆炸
# Pre-Norm(现代): x + Sublayer(LN(x)) ← 训练更稳定,收敛更快
# 几乎所有现代LLM(GPT/LLaMA/Qwen)都使用Pre-Norm
七、组装完整GPT-2模型
python
class GPT2Mini(nn.Module):
"""迷你GPT-2: 可训练的语言模型"""
def __init__(self, vocab_size=50257, d_model=512, n_heads=8,
n_layers=12, d_ff=2048, max_len=1024, dropout=0.1):
super().__init__()
# Token嵌入
self.token_embedding = nn.Embedding(vocab_size, d_model)
# 位置嵌入(GPT-2使用可学习的位置嵌入,非sin/cos)
self.position_embedding = nn.Embedding(max_len, d_model)
self.dropout = nn.Dropout(dropout)
# Transformer层堆叠
self.layers = nn.ModuleList([
TransformerBlock(d_model, n_heads, d_ff, dropout)
for _ in range(n_layers)
])
# 最终LayerNorm
self.ln_f = nn.LayerNorm(d_model)
# 输出头(预测下一个token)
self.lm_head = nn.Linear(d_model, vocab_size, bias=False)
# ★ 权重绑定:输入嵌入和输出头共享权重
# 这减少了参数量,且有助于正则化
self.lm_head.weight = self.token_embedding.weight
# 初始化
self.apply(self._init_weights)
def _init_weights(self, module):
if isinstance(module, nn.Linear):
torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
if module.bias is not None:
torch.nn.init.zeros_(module.bias)
elif isinstance(module, nn.Embedding):
torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
def forward(self, input_ids, attention_mask=None):
"""
input_ids: [batch_size, seq_len]
attention_mask: [batch_size, seq_len] --- 1表示有效token,0表示padding
"""
batch_size, seq_len = input_ids.shape
device = input_ids.device
# 1. Token嵌入 + 位置嵌入
token_embeds = self.token_embedding(input_ids)
pos_ids = torch.arange(0, seq_len, dtype=torch.long, device=device)
pos_embeds = self.position_embedding(pos_ids)
# 2. 合并嵌入
x = self.dropout(token_embeds + pos_embeds) # [B, T, C]
# 3. 构造因果mask + padding mask
# 因果mask: 下三角矩阵
causal_mask = torch.tril(
torch.ones(seq_len, seq_len, device=device)
).view(1, 1, seq_len, seq_len)
# Padding mask: 将padding位置也mask掉
if attention_mask is not None:
# attention_mask: [B, T] → [B, 1, 1, T]
padding_mask = attention_mask[:, None, None, :]
# 结合因果mask
combined_mask = causal_mask * padding_mask
else:
combined_mask = causal_mask
# 4. Transformer层前向传播
all_attentions = []
for layer in self.layers:
x, attn_weights = layer(x, combined_mask)
all_attentions.append(attn_weights)
# 5. 最终归一化 + 输出投影
x = self.ln_f(x)
logits = self.lm_head(x) # [B, T, vocab_size]
return logits, all_attentions
def generate(self, input_ids, max_new_tokens=50, temperature=0.8, top_k=40):
"""自回归文本生成"""
self.eval()
for _ in range(max_new_tokens):
# 截断到最大长度(处理超长序列)
if input_ids.size(1) > 1024:
input_ids = input_ids[:, -1024:]
# 前向传播
with torch.no_grad():
logits, _ = self(input_ids)
# 取最后一个位置的logits
logits = logits[:, -1, :] / temperature
# Top-K采样
if top_k > 0:
indices_to_remove = logits < torch.topk(logits, top_k)[0][..., -1, None]
logits[indices_to_remove] = float('-inf')
# 采样
probs = F.softmax(logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
# 拼接
input_ids = torch.cat([input_ids, next_token], dim=-1)
return input_ids
# ============ 训练示例 ============
def train_gpt2_mini():
from torch.utils.data import DataLoader
from transformers import GPT2Tokenizer
tokenizer = GPT2Tokenizer.from_pretrained("gpt2")
tokenizer.pad_token = tokenizer.eos_token
model = GPT2Mini(
vocab_size=tokenizer.vocab_size,
d_model=512,
n_heads=8,
n_layers=6, # 6层轻量版
d_ff=2048,
max_len=512
)
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=0.01)
# 模拟训练数据(实际应用中替换为真实语料)
texts = [
"The quick brown fox jumps over the lazy dog",
"Machine learning is a subset of artificial intelligence",
# ... 更多训练数据
]
for epoch in range(3):
total_loss = 0
for text in texts:
# Tokenize
tokens = tokenizer.encode(text)
input_ids = torch.tensor([tokens])
# 前向传播
logits, _ = model(input_ids)
# 损失:预测下一个token(shift left)
shift_logits = logits[:, :-1, :].contiguous()
shift_labels = input_ids[:, 1:].contiguous()
loss = F.cross_entropy(
shift_logits.view(-1, shift_logits.size(-1)),
shift_labels.view(-1)
)
optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
total_loss += loss.item()
print(f"Epoch {epoch+1}, Loss: {total_loss/len(texts):.4f}")
# 生成测试
prompt = tokenizer.encode("The future of AI is")
input_ids = torch.tensor([prompt])
generated = model.generate(input_ids, max_new_tokens=30)
print(f"\nGenerated: {tokenizer.decode(generated[0])}")
# 参数量估算:
# Embedding: vocab(50257) × d_model(512) ≈ 25.7M
# Per Layer: 4×(512²) + 2×(512×2048) ≈ 3.1M
# 6 Layers: ≈ 18.6M
# Total GPT2Mini(6层): ≈ 44M参数
# Total GPT2 Small(12层): ≈ 124M参数
八、进化:RoPE旋转位置编码
python
class RotaryPositionalEmbedding(nn.Module):
"""RoPE: 通过在Q和K上施加旋转矩阵来编码相对位置"""
def __init__(self, dim, max_position_embeddings=2048, base=10000):
super().__init__()
self.dim = dim
self.max_position_embeddings = max_position_embeddings
self.base = base
# 预计算旋转频率矩阵
inv_freq = 1.0 / (base ** (
torch.arange(0, dim, 2).float() / dim
))
self.register_buffer("inv_freq", inv_freq)
# 缓存cos和sin值(加速推理)
self._set_cos_sin_cache(max_position_embeddings)
def _set_cos_sin_cache(self, seq_len):
t = torch.arange(seq_len).type_as(self.inv_freq)
freqs = torch.outer(t, self.inv_freq) # [seq_len, dim/2]
emb = torch.cat((freqs, freqs), dim=-1) # [seq_len, dim]
self.register_buffer("cos_cached", emb.cos()[None, None, :, :])
self.register_buffer("sin_cached", emb.sin()[None, None, :, :])
def rotate_half(self, x):
"""将向量的后半部分旋转到前半部分(90度旋转)"""
x1, x2 = x[..., :x.shape[-1]//2], x[..., x.shape[-1]//2:]
return torch.cat((-x2, x1), dim=-1)
def apply_rotary_pos_emb(self, q, k, cos, sin):
"""对Q和K应用旋转位置编码"""
# q, k: [batch, heads, seq, dim]
cos = cos[:, :, :q.size(2), :]
sin = sin[:, :, :q.size(2), :]
q_embed = (q * cos) + (self.rotate_half(q) * sin)
k_embed = (k * cos) + (self.rotate_half(k) * sin)
return q_embed, k_embed
# RoPE的优势:
# 1. 相对位置: Q_m·K_n = f(m-n),自然编码相对距离
# 2. 远程衰减: |m-n|越大,内积越小(自然的距离衰减)
# 3. 无需训练: 与sin/cos一样是固定的
# 4. 被LLaMA/Qwen/Mistral广泛采用
九、总结
从零实现了完整的Transformer全链路:
- Scaled Dot-Product Attention :
softmax(QK^T/√d_k)V,除以√d_k控制方差 - Multi-Head Attention: 8个并行头,每个头关注不同特征子空间
- Positional Encoding: sin/cos编码相对位置,RoPE在现代模型中广泛使用
- FFN + LayerNorm + 残差: 非线性变换 + Pre-Norm稳定训练
- 完整GPT-2: 12层堆叠 + 自回归生成 + Top-K采样
理解了这些基础后,GPT-4/Claude/Llama 等模型的核心思想都在此------只是更深(数百层)、更宽(万维隐藏层)、更多数据(万亿token)。