从0搭建极简transformer大模型
已在 cucode/trans.py 中创建了一个 基于 PyTorch 的 Transformer 模型从零实现,共 797 行代码,所有测试和训练均验证通过。
Transformer | 如何理解 attention 中的 Q,K,V?
Q、K、V 与 Multi-Head Attention 多头注意力机制
仅仅使用pytorch来手撕transformer架构(2):多头注意力MultiHeadAttention类的实现和向前传播
文件结构
| 模块 | 内容 |
|---|---|
| 核心组件 | attention() 缩放点积注意力、MultiHeadAttention 多头注意力、PositionalEncoding 正弦位置编码、PositionwiseFeedForward 前馈网络、LayerNorm 层归一化、SublayerConnection 残差连接 |
| 编码器 | EncoderLayer(自注意力 + FFN + 残差归一化)、Encoder(N 层堆叠) |
| 解码器 | DecoderLayer(掩码自注意力 + 交叉注意力 + FFN)、Decoder(N 层堆叠) |
| 完整模型 | Transformer(Encoder-Decoder 架构)、Embeddings 词嵌入、Generator 输出层 |
| 工具函数 | make_model() 模型构建、subsequent_mask() 因果掩码 |
| 训练示例 | CharDataset 字符级数据集、train_language_model() 训练流程、generate_text() 自回归生成 |
| 单元测试 | test_components() 验证所有组件的形状和梯度正确性 |
运行结果
组件测试: 全部 8 项通过
- Scaled Dot-Product Attention
- Multi-Head Attention
- Positional Encoding
- Feed-Forward Network
- LayerNorm
- Subsequent Mask
- Full Transformer Model
- Backpropagation
训练示例:
Epoch 1/30 | Loss: 1.7620
Epoch 30/30 | Loss: 0.0150 (损失下降 99%+)
生成: 'hello world,lethisisisis...' (能复现训练文本模式)
关键设计特点
- 忠实论文:严格按 “Attention Is All You Need” 原文实现,包括 Post-LN 结构、Xavier 初始化、正弦位置编码
- 详细注释:每个类和函数都有中文文档字符串,解释公式、参数含义和维度变换
- 可配置:
make_model()支持自定义层数、维度、头数等超参数 - 即开即用:
python cucode/trans.py一键运行测试 + 训练演示
"""
"""
trans.py
========
基于 PyTorch 2.x 的 Transformer 模型从零实现(学习用途)。
本文件实现了论文 "Attention Is All You Need" (Vaswani et al., 2017) 中
提出的完整 Transformer 架构,包括:
1. 多头自注意力机制 (Multi-Head Self-Attention)
2. 位置编码 (Positional Encoding)
3. 前馈网络 (Position-wise Feed-Forward Network)
4. 编码器层 / 编码器 (Encoder Layer / Encoder)
5. 解码器层 / 解码器 (Decoder Layer / Decoder)
6. 完整的 Transformer 模型 (seq2seq)
7. 一个简单的语言模型训练示例(字符级)
运行环境: Python 3.10+, PyTorch 2.x
依赖: pip install torch
"""
import math
import copy
import time
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
# ============================================================================
# 第一部分:模型核心组件
# ============================================================================
def clones(module: nn.Module, n: int) -> nn.ModuleList:
"""复制 n 个完全相同(但参数独立)的子模块。"""
return nn.ModuleList([copy.deepcopy(module) for _ in range(n)])
def attention(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
mask: torch.Tensor | None = None,
dropout: nn.Dropout | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
缩放点积注意力 (Scaled Dot-Product Attention)。
公式: Attention(Q, K, V) = softmax(Q @ K^T / sqrt(d_k)) @ V
参数:
query : (batch, n_heads, seq_len, d_k)
key : (batch, n_heads, seq_len, d_k)
value : (batch, n_heads, seq_len, d_v)
mask : (batch, 1, 1, seq_len) 或 None,被屏蔽位置设为 0
dropout: 可选的 dropout 层
返回:
output: 注意力加权后的值 (batch, n_heads, seq_len, d_v)
p_attn: 注意力权重矩阵 (batch, n_heads, seq_len, seq_len)
"""
d_k = query.size(-1)
# 1. 计算注意力分数: Q @ K^T / sqrt(d_k)
scores = torch.matmul(query, key.transpose(-2, -1)) / math.sqrt(d_k)
# 2. 应用 mask(将不该看到的位置分数设为 -inf,softmax 后趋近于 0)
if mask is not None:
scores = scores.masked_fill(mask == 0, float("-inf"))
# 3. softmax 归一化得到注意力权重
p_attn = F.softmax(scores, dim=-1)
# 4. 可选 dropout
if dropout is not None:
p_attn = dropout(p_attn)
# 5. 用注意力权重对 value 加权求和
output = torch.matmul(p_attn, value)
return output, p_attn
class MultiHeadAttention(nn.Module):
"""
多头注意力机制 (Multi-Head Attention)。
将 Q、K、V 分别投影到 h 个不同的子空间,各自做缩放点积注意力,
最后拼接所有头的输出并做一次线性投影。
参数:
n_heads: 注意力头数
d_model: 模型维度(必须能被 n_heads 整除)
dropout: dropout 比率
"""
def __init__(self, n_heads: int, d_model: int, dropout: float = 0.1):
super().__init__()
assert d_model % n_heads == 0, "d_model 必须能被 n_heads 整除"
self.d_k = d_model // n_heads # 每个头的维度
self.n_heads = n_heads
# 4 个线性层: Q, K, V 的投影 + 输出投影
self.linears = clones(nn.Linear(d_model, d_model), 4)
self.attn: torch.Tensor | None = None # 保存注意力权重(用于可视化/调试)
self.dropout = nn.Dropout(p=dropout)
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
mask: torch.Tensor | None = None,
) -> torch.Tensor:
"""
参数:
query, key, value: (batch, seq_len, d_model)
mask: (batch, 1, seq_len) 或 None
返回:
(batch, seq_len, d_model)
"""
if mask is not None:
# mask 形状: (batch, 1, 1, seq_len),方便广播到所有头
mask = mask.unsqueeze(1)
n_batches = query.size(0)
# 1. 对 Q, K, V 做线性投影,并 reshape 为多头形式
# (batch, seq_len, d_model) -> (batch, n_heads, seq_len, d_k)
query, key, value = [
lin(x).view(n_batches, -1, self.n_heads, self.d_k).transpose(1, 2)
for lin, x in zip(self.linears, (query, key, value))
]
# 2. 计算注意力
out, self.attn = attention(query, key, value, mask=mask, dropout=self.dropout)
# 3. 拼接所有头: (batch, n_heads, seq_len, d_k) -> (batch, seq_len, d_model)
out = out.transpose(1, 2).contiguous().view(n_batches, -1, self.n_heads * self.d_k)
# 4. 最终线性投影
return self.linears[-1](out)
class PositionwiseFeedForward(nn.Module):
"""
位置前馈网络 (Position-wise Feed-Forward Network)。
对每个位置独立地做两层线性变换 + ReLU 激活:
FFN(x) = W2 * ReLU(W1 * x + b1) + b2
参数:
d_model: 模型输入/输出维度
d_ff: 中间隐藏层维度(通常为 4 * d_model)
dropout: dropout 比率
"""
def __init__(self, d_model: int, d_ff: int, dropout: float = 0.1):
super().__init__()
self.w_1 = nn.Linear(d_model, d_ff)
self.w_2 = nn.Linear(d_ff, d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.w_2(self.dropout(F.relu(self.w_1(x))))
class PositionalEncoding(nn.Module):
"""
正弦位置编码 (Sinusoidal Positional Encoding)。
为序列中每个位置生成一个固定的位置向量,使用不同频率的正弦/余弦函数:
PE(pos, 2i) = sin(pos / 10000^(2i/d_model))
PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))
位置编码与词嵌入相加,使模型能感知序列中 token 的顺序。
参数:
d_model: 模型维度
dropout: dropout 比率
max_len: 预计算的最大序列长度
"""
def __init__(self, d_model: int, dropout: float, max_len: int = 5000):
super().__init__()
self.dropout = nn.Dropout(p=dropout)
# (max_len, 1) 位置索引
position = torch.arange(0, max_len).unsqueeze(1).float()
# (d_model // 2,) 频率项: 10000^(2i/d_model)
div_term = torch.exp(
torch.arange(0, d_model, 2).float() * -(math.log(10000.0) / d_model)
)
pe = torch.zeros(max_len, d_model) # (max_len, d_model)
pe[:, 0::2] = torch.sin(position * div_term) # 偶数位用 sin
pe[:, 1::2] = torch.cos(position * div_term) # 奇数位用 cos
# 增加 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
"""
x = x + self.pe[:, : x.size(1)]
return self.dropout(x)
class LayerNorm(nn.Module):
"""
层归一化 (Layer Normalization)。
对最后一个维度做归一化: y = (x - mean) / sqrt(var + eps) * gamma + beta
参数:
features: 归一化的特征维度
eps: 防止除以零的小常数
"""
def __init__(self, features: int, eps: float = 1e-6):
super().__init__()
self.gamma = nn.Parameter(torch.ones(features)) # 可学习的缩放
self.beta = nn.Parameter(torch.zeros(features)) # 可学习的偏移
self.eps = eps
def forward(self, x: torch.Tensor) -> torch.Tensor:
mean = x.mean(-1, keepdim=True)
std = x.std(-1, keepdim=True)
return self.gamma * (x - mean) / (std + self.eps) + self.beta
class SublayerConnection(nn.Module):
"""
残差连接 + LayerNorm (Sublayer Connection)。
采用 "Post-LN" 结构(原论文写法):
output = LayerNorm(x + Sublayer(x))
参数:
size: 特征维度
dropout: dropout 比率
"""
def __init__(self, size: int, dropout: float):
super().__init__()
self.norm = LayerNorm(size)
self.dropout = nn.Dropout(dropout)
def forward(self, x: torch.Tensor, sublayer: nn.Module) -> torch.Tensor:
"""将子层的输出经过 dropout 后与输入相加,再做 LayerNorm。"""
return self.norm(x + self.dropout(sublayer(x)))
# ============================================================================
# 第二部分:编码器与解码器
# ============================================================================
class EncoderLayer(nn.Module):
"""
编码器层 (Encoder Layer)。
每层包含两个子层:
1. 多头自注意力 (Self-Attention)
2. 前馈网络 (Feed-Forward)
每个子层都配有残差连接 + LayerNorm。
参数:
size: 模型维度
attn: 多头注意力模块
ff: 前馈网络模块
dropout: dropout 比率
"""
def __init__(self, size: int, attn: MultiHeadAttention, ff: PositionwiseFeedForward, dropout: float):
super().__init__()
self.size = size
self.self_attn = attn
self.ff = ff
# 两个 SublayerConnection 分别用于注意力子层和前馈子层
self.sublayer = clones(SublayerConnection(size, dropout), 2)
def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
"""参数: x (batch, seq_len, d_model), mask (batch, 1, seq_len)"""
# 子层1: 自注意力 (Q=K=V=x)
x = self.sublayer[0](x, lambda t: self.self_attn(t, t, t, mask))
# 子层2: 前馈网络
x = self.sublayer[1](x, self.ff)
return x
class Encoder(nn.Module):
"""将 N 个 EncoderLayer 堆叠起来。"""
def __init__(self, layer: EncoderLayer, n: int):
super().__init__()
self.layers = clones(layer, n)
self.norm = LayerNorm(layer.size) # 最后的整体 LayerNorm
def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
for layer in self.layers:
x = layer(x, mask)
return self.norm(x)
class DecoderLayer(nn.Module):
"""
解码器层 (Decoder Layer)。
每层包含三个子层:
1. 掩码多头自注意力 (Masked Self-Attention)
2. 编码器-解码器交叉注意力 (Cross-Attention)
3. 前馈网络 (Feed-Forward)
每个子层都配有残差连接 + LayerNorm。
参数:
size: 模型维度
attn: 多头注意力模块(用于自注意力和交叉注意力)
ff: 前馈网络模块
dropout: dropout 比率
"""
def __init__(self, size: int, attn: MultiHeadAttention, ff: PositionwiseFeedForward, dropout: float):
super().__init__()
self.size = size
self.self_attn = copy.deepcopy(attn) # 自注意力
self.cross_attn = copy.deepcopy(attn) # 交叉注意力
self.ff = ff
self.sublayer = clones(SublayerConnection(size, dropout), 3)
def forward(
self,
x: torch.Tensor,
memory: torch.Tensor,
src_mask: torch.Tensor,
tgt_mask: torch.Tensor,
) -> torch.Tensor:
"""
参数:
x: 解码器输入 (batch, tgt_len, d_model)
memory: 编码器输出 (batch, src_len, d_model)
src_mask: 源序列 padding mask
tgt_mask: 目标序列 mask(含因果掩码 + padding mask)
"""
# 子层1: 掩码自注意力
x = self.sublayer[0](x, lambda t: self.self_attn(t, t, t, tgt_mask))
# 子层2: 交叉注意力 (Q=x, K=V=memory)
x = self.sublayer[1](x, lambda t: self.cross_attn(t, memory, memory, src_mask))
# 子层3: 前馈网络
x = self.sublayer[2](x, self.ff)
return x
class Decoder(nn.Module):
"""将 N 个 DecoderLayer 堆叠起来。"""
def __init__(self, layer: DecoderLayer, n: int):
super().__init__()
self.layers = clones(layer, n)
self.norm = LayerNorm(layer.size)
def forward(
self,
x: torch.Tensor,
memory: torch.Tensor,
src_mask: torch.Tensor,
tgt_mask: torch.Tensor,
) -> torch.Tensor:
for layer in self.layers:
x = layer(x, memory, src_mask, tgt_mask)
return self.norm(x)
# ============================================================================
# 第三部分:嵌入层、生成器与完整模型
# ============================================================================
class Embeddings(nn.Module):
"""词嵌入层: 将 token id 映射为 d_model 维向量。"""
def __init__(self, d_model: int, vocab: int):
super().__init__()
self.lut = nn.Embedding(vocab, d_model)
self.d_model = d_model
def forward(self, x: torch.Tensor) -> torch.Tensor:
# 乘以 sqrt(d_model) 使嵌入值与位置编码量级匹配
return self.lut(x) * math.sqrt(self.d_model)
class Generator(nn.Module):
"""输出层: 线性投影 + log-softmax,将隐藏状态映射到词表上的概率分布。"""
def __init__(self, d_model: int, vocab: int):
super().__init__()
self.proj = nn.Linear(d_model, vocab)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return F.log_softmax(self.proj(x), dim=-1)
class Transformer(nn.Module):
"""
完整的 Encoder-Decoder Transformer 模型。
参数:
encoder: 编码器
decoder: 解码器
src_embed: 源序列嵌入(词嵌入 + 位置编码)
tgt_embed: 目标序列嵌入(词嵌入 + 位置编码)
generator: 输出生成器
"""
def __init__(self, encoder: Encoder, decoder: Decoder, src_embed: nn.Module, tgt_embed: nn.Module, generator: Generator):
super().__init__()
self.encoder = encoder
self.decoder = decoder
self.src_embed = src_embed
self.tgt_embed = tgt_embed
self.generator = generator
def forward(self, src: torch.Tensor, tgt: torch.Tensor, src_mask: torch.Tensor, tgt_mask: torch.Tensor) -> torch.Tensor:
"""
参数:
src: 源序列 token ids (batch, src_len)
tgt: 目标序列 token ids (batch, tgt_len)
src_mask: 源序列 padding mask
tgt_mask: 目标序列 mask(因果 + padding)
返回:
log 概率 (batch, tgt_len, vocab_size)
"""
return self.decode(
self.encode(src, src_mask), src_mask, tgt, tgt_mask
)
def encode(self, src: torch.Tensor, src_mask: torch.Tensor) -> torch.Tensor:
return self.encoder(self.src_embed(src), src_mask)
def decode(self, memory: torch.Tensor, src_mask: torch.Tensor, tgt: torch.Tensor, tgt_mask: torch.Tensor) -> torch.Tensor:
return self.generator(self.decoder(self.tgt_embed(tgt), memory, src_mask, tgt_mask))
# ============================================================================
# 第四部分:模型构建函数
# ============================================================================
def make_model(
src_vocab: int,
tgt_vocab: int,
n_layers: int = 6,
d_model: int = 512,
d_ff: int = 2048,
n_heads: int = 8,
dropout: float = 0.1,
) -> Transformer:
"""
构建并初始化一个完整的 Transformer 模型。
参数:
src_vocab: 源语言词表大小
tgt_vocab: 目标语言词表大小
n_layers: 编码器/解码器层数 (默认 6)
d_model: 模型维度 (默认 512)
d_ff: 前馈网络中间维度 (默认 2048)
n_heads: 注意力头数 (默认 8)
dropout: dropout 比率 (默认 0.1)
返回:
Transformer 模型实例
"""
# 创建模型组件
attn = MultiHeadAttention(n_heads, d_model, dropout)
ff = PositionwiseFeedForward(d_model, d_ff, dropout)
position = PositionalEncoding(d_model, dropout)
model = Transformer(
encoder=Encoder(EncoderLayer(d_model, copy.deepcopy(attn), copy.deepcopy(ff), dropout), n_layers),
decoder=Decoder(DecoderLayer(d_model, copy.deepcopy(attn), copy.deepcopy(ff), dropout), n_layers),
src_embed=nn.Sequential(Embeddings(d_model, src_vocab), copy.deepcopy(position)),
tgt_embed=nn.Sequential(Embeddings(d_model, tgt_vocab), copy.deepcopy(position)),
generator=Generator(d_model, tgt_vocab),
)
# 参数初始化: Xavier 均匀分布
for p in model.parameters():
if p.dim() > 1:
nn.init.xavier_uniform_(p)
print(f"模型构建完成 | 参数量: {sum(p.numel() for p in model.parameters()):,}")
return model
# ============================================================================
# 第五部分:Mask 工具函数
# ============================================================================
def subsequent_mask(size: int) -> torch.Tensor:
"""
生成因果掩码 (Causal Mask),防止解码器"看到"未来位置。
返回下三角矩阵: (1, size, size),上三角部分为 0(被屏蔽),下三角为 1。
示例 (size=4):
[[1, 0, 0, 0],
[1, 1, 0, 0],
[1, 1, 1, 0],
[1, 1, 1, 1]]
"""
attn_shape = (1, size, size)
mask = torch.triu(torch.ones(attn_shape), diagonal=1).type(torch.uint8)
return mask == 0 # True 的位置允许注意
# ============================================================================
# 第六部分:简单的字符级语言模型训练示例
# ============================================================================
class CharDataset(Dataset):
"""
字符级数据集: 从一段文本中生成训练样本。
采用"前缀续写"任务(encoder-decoder 架构的正确用法):
prefix = 文本块的前半部分 -> 作为 encoder 输入("全局前缀")
x = 完整文本块 -> 作为 decoder 输入(自回归)
y = x 右移一位 -> decoder 的目标(预测下一个字符)
关键设计: encoder 在训练时只看到前缀(看不到未来),与推理时一致,
避免"encoder 偷看未来导致训练/推理分布不匹配"的问题。
"""
def __init__(self, text: str, seq_len: int = 48):
self.chars = sorted(set(text))
self.char2idx = {c: i for i, c in enumerate(self.chars)}
self.idx2char = {i: c for c, i in self.char2idx.items()}
self.vocab_size = len(self.chars)
self.seq_len = seq_len
# encoder 看到的前缀长度(固定为序列前一半,训练与推理一致)
self.prefix_len = seq_len // 2
# 将文本编码为 token id 序列
self.data = torch.tensor([self.char2idx[c] for c in text], dtype=torch.long)
def __len__(self) -> int:
return max(0, (len(self.data) - self.seq_len - 1) // 1)
def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
chunk = self.data[idx : idx + self.seq_len + 1]
prefix = chunk[: self.prefix_len] # encoder 输入(前一半,无未来信息)
x = chunk[:-1] # decoder 输入
y = chunk[1:] # decoder 目标(右移一位)
return prefix, x, y
def train_language_model():
"""
训练"前缀续写"任务,演示 Encoder-Decoder Transformer 的完整训练流程。
与 decoder_only.py(纯 GPT 式自回归)不同,这里同时使用编码器和解码器:
encoder 读取提示前缀 -> decoder 基于前缀自回归地续写文本
这种"前缀续写"训练方式避免了经典 encoder-decoder 语言模型中
"encoder 双向看到未来、推理时却看不到"导致的训练/推理分布不匹配。
"""
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"训练设备: {device}")
# --- 1. 准备数据 ---
# 用于学习的示例文本:少而规律的句子(与 decoder_only.py 类似,
# 让模型能精确学会字符间的转移规律),配合 n-gram 阻断采样防复读
sample_text = (
"the quick brown fox jumps over the lazy dog. "
"pack my box with five dozen liquor jugs. "
"attention is all you need, said the wise model. "
"deep learning is fun and powerful to study. "
"let us build a small language model together. "
"practice makes perfect, keep coding every day. "
) * 30 # 重复以增加训练数据量
seq_len = 48 # 上下文长度(decoder 能看到的字符数)
dataset = CharDataset(sample_text, seq_len=seq_len)
dataloader = DataLoader(dataset, batch_size=16, shuffle=True)
vocab_size = dataset.vocab_size
print(f"词表大小: {vocab_size}, 字符: {''.join(dataset.chars)}")
print(f"训练样本数: {len(dataset)}")
# --- 2. 构建模型 ---
# 对于自回归语言模型,src 和 tgt 共享同一个词表
model = make_model(
src_vocab=vocab_size,
tgt_vocab=vocab_size,
n_layers=2, # 小模型: 2 层
d_model=64, # 小维度
d_ff=256,
n_heads=4,
dropout=0.2, # 提高 dropout 抑制过拟合
).to(device)
# --- 3. 优化器 & 损失函数 ---
# AdamW + 权重衰减,训练更稳定(GPT 系列标准做法)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01)
criterion = nn.NLLLoss(ignore_index=0) # 负对数似然损失
# 学习率调度: warmup + 余弦衰减(warmup 阶段逐步提高 lr,之后余弦降到 0)
epochs = 30 # 训练轮数:让模型充分学习这 8 个规律句子
total_steps = len(dataloader) * epochs
warmup_steps = max(1, int(0.1 * total_steps))
def lr_lambda(step: int) -> float:
if step < warmup_steps:
return step / warmup_steps # 线性 warmup
progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)
return 0.5 * (1.0 + math.cos(math.pi * progress)) # 余弦衰减
scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)
# --- 4. 训练循环 ---
model.train()
print("\n开始训练...")
print("=" * 60)
for epoch in range(epochs):
t0 = time.time()
total_loss = 0.0
n_batches = 0
for prefix, x, y in dataloader:
prefix, x, y = prefix.to(device), x.to(device), y.to(device)
# 构造 mask:
# src_mask: encoder 输入(前缀)全部有效
# tgt_mask: decoder 输入使用因果掩码(防止看到未来)
src_mask = torch.ones(prefix.size(0), 1, prefix.size(1), device=device)
tgt_mask = subsequent_mask(x.size(1)).to(device) # (1, seq_len, seq_len)
tgt_mask = tgt_mask.expand(x.size(0), -1, -1) # (batch, seq_len, seq_len)
# 前向传播: encoder 读前缀,decoder 自回归续写
logits = model(prefix, x, src_mask, tgt_mask) # (batch, seq_len, vocab)
# 计算损失: 展平后与目标比较
loss = criterion(
logits.reshape(-1, vocab_size),
y.reshape(-1),
)
# 反向传播
optimizer.zero_grad()
loss.backward()
optimizer.step()
scheduler.step() # 更新学习率
total_loss += loss.item()
n_batches += 1
avg_loss = total_loss / max(n_batches, 1)
elapsed = time.time() - t0
if (epoch + 1) % 5 == 0 or epoch == 0:
print(f"Epoch {epoch + 1:3d}/{epochs} | Loss: {avg_loss:.4f} | Time: {elapsed:.2f}s")
print("=" * 60)
print("训练完成!")
# --- 5. 生成文本 ---
print("\n--- 文本生成示例 ---")
model.eval()
# 提示词取自训练文本,长度对齐训练时的 prefix_len=24,保证训练/推理分布一致
prompts = [
"the quick brown fox jumps o", # 24 字符
"pack my box with five doze", # 24 字符
"attention is all you need, ", # 24 字符
"deep learning is fun and po", # 24 字符
]
for prompt in prompts:
# 贪心解码 + n-gram 阻断:
# temperature=0 退化为贪心(取概率最大),可精确复现训练文本;
# n-gram 阻断负责防止陷入"复读机"死循环(如 sisisisi...)
generated = generate_text(
model, prompt, dataset,
length=32, device=device,
temperature=0.0, top_k=None, top_p=None,
repetition_penalty=1.0, no_repeat_ngram_size=3,
)
print(f"提示: {prompt!r}")
print(f"生成: {generated!r}")
print()
return model, dataset
def generate_text(
model: Transformer,
prompt: str,
dataset: CharDataset,
length: int = 40,
device: torch.device = torch.device("cpu"),
temperature: float = 0.8,
top_k: int | None = 8,
top_p: float | None = 0.9,
repetition_penalty: float = 2.0,
no_repeat_ngram_size: int = 3,
) -> str:
"""
前缀续写文本生成(Encoder-Decoder 架构的正确用法)。
流程:
1. prompt 作为"前缀"输入 encoder(编码全局上下文,固定不变)
2. decoder 从 prompt 开始,每次取最后一个位置的概率分布采样一个
新字符追加到序列末尾
3. 重复第 2 步 length 次
这样训练(encoder 只看前缀)与推理(encoder 只看 prompt)的分布一致,
避免了"encoder 双向偷看未来"导致的训练/推理不匹配。
相比直接取最大概率的贪心解码(argmax),这里使用多种采样策略来
避免模型陷入"复读机"死循环:
temperature: 温度缩放。>1 更随机,<1 更保守,0 退化为贪心
top_k: Top-K 采样,只从前 K 个概率最高的 token 中采样
top_p: Top-P (Nucleus) 采样,从累积概率达 p 的最小集合采样
repetition_penalty: 重复惩罚,对已出现过的 token 概率打折,抑制重复
no_repeat_ngram_size: n-gram 阻断,禁止生成会形成已出现过 n-gram
的 token,是防止复读的最强手段(默认 3)
参数:
model: 训练好的 Transformer 模型
prompt: 提示文本(将作为 encoder 输入的前缀)
dataset: 数据集(用于字符到 id 的映射)
length: 要生成的字符数
device: 计算设备
temperature: 采样温度。>1 更随机,<1 更保守,0 或负数退化为贪心解码
(配合 no_repeat_ngram_size 使用可实现"精确复现 + 防复读")
top_k: 保留概率最高的前 k 个 token(默认 8,None 表示禁用)
top_p: 保留累积概率前 p 的 token 集合(默认 0.9,None 表示禁用)
repetition_penalty: 已出现 token 的惩罚系数(默认 2.0,1.0 表示禁用)
no_repeat_ngram_size: n-gram 阻断窗口大小(默认 3,<=1 表示禁用)
"""
model.eval()
# 将 prompt 编码为 token ids
ids = [dataset.char2idx.get(c, 0) for c in prompt]
# encoder 输入 = prompt(前缀,固定不变)
src = torch.tensor([ids], dtype=torch.long, device=device) # (1, prompt_len)
src_mask = torch.ones(1, 1, src.size(1), device=device) # (1, 1, prompt_len)
# decoder 初始输入 = prompt(从 prompt 开始续写)
tgt = torch.tensor([ids], dtype=torch.long, device=device) # (1, prompt_len)
def _no_repeat_ngram_block(logits: torch.Tensor) -> torch.Tensor:
"""n-gram 阻断: 若生成的 token 会形成已生成序列中已出现过的 n-gram,则禁止之。"""
if no_repeat_ngram_size <= 1:
return logits
seq = tgt[0].tolist()
n = no_repeat_ngram_size
if len(seq) < n:
return logits
# 当前前缀 = 最后 n-1 个 token
prefix = tuple(seq[-(n - 1):])
# 扫描整个序列,收集"该前缀之后出现过哪些 token"
banned: set[int] = set()
for i in range(len(seq) - n + 1):
window = tuple(seq[i:i + n])
if window[:n - 1] == prefix:
banned.add(window[-1])
# 将这些 token 的分数设为 -inf(禁止生成)
if banned:
for token_id in banned:
logits[token_id] = float("-inf")
return logits
with torch.no_grad():
for _ in range(length):
tgt_len = tgt.size(1)
tgt_mask = subsequent_mask(tgt_len).to(device) # (1, tgt_len, tgt_len)
# 前向传播: encoder 编码前缀,decoder 基于前缀自回归续写
logits = model(src, tgt, src_mask, tgt_mask) # (1, tgt_len, vocab)
# 取最后一个位置的原始分数 (vocab,)
next_logits = logits[0, -1, :].clone()
# --- 采样策略 ---
# 1. 重复惩罚: 对已生成序列中出现过的 token 的分数打折
# 分数越高的 token 受影响越大,从而抑制模型重复输出
if repetition_penalty > 1.0:
for token_id in set(tgt[0].tolist()):
if next_logits[token_id] > 0:
next_logits[token_id] /= repetition_penalty
else:
next_logits[token_id] *= repetition_penalty
# 2. n-gram 阻断: 防止生成已出现过的 n-gram(抑制复读的最强手段)
next_logits = _no_repeat_ngram_block(next_logits)
# 3. 温度缩放: 除以温度后,softmax 分布的"锐度"发生变化
# 当 temperature <= 0 时退化为贪心解码(直接取概率最大的 token)
if temperature <= 0.0:
next_id = next_logits.argmax().item()
tgt = torch.cat(
[tgt, torch.tensor([[next_id]], dtype=torch.long, device=device)],
dim=1,
)
continue # 跳过后续采样步骤,直接进入下一轮
if temperature != 1.0:
next_logits = next_logits / temperature
# 4. Top-K: 只保留概率最高的前 k 个 token,其余设为 -inf
if top_k is not None:
v, _ = torch.topk(next_logits, min(top_k, next_logits.size(-1)))
next_logits[next_logits < v[-1]] = float("-inf")
# 5. Top-P (Nucleus): 只保留累积概率达到 top_p 的最小 token 集合
if top_p is not None and 0.0 < top_p < 1.0:
sorted_logits, sorted_indices = torch.sort(next_logits, descending=True)
cum_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
# 移除累积概率超过 top_p 的 token(保留第一个,避免全被移除)
sorted_mask = cum_probs > top_p
sorted_mask[..., 1:] = sorted_mask[..., :-1].clone()
sorted_mask[..., 0] = False
# 映射回原始索引
indices_to_remove = sorted_mask.scatter(-1, sorted_indices, sorted_mask)
next_logits = next_logits.masked_fill(indices_to_remove, float("-inf"))
# 6. 从概率分布中随机采样一个 token
probs = F.softmax(next_logits, dim=-1)
next_id = torch.multinomial(probs, num_samples=1).item()
# 追加到 decoder 序列
tgt = torch.cat(
[tgt, torch.tensor([[next_id]], dtype=torch.long, device=device)],
dim=1,
)
# 解码为文本
result_ids = tgt[0].cpu().tolist()
result = "".join(dataset.idx2char.get(i, "?") for i in result_ids)
return result
# ============================================================================
# 第七部分:单元测试(验证模型各组件的正确性)
# ============================================================================
def test_components():
"""运行一系列断言测试,验证模型各组件的输入输出形状是否正确。"""
print("\n" + "=" * 60)
print("运行组件测试...")
print("=" * 60)
batch, seq_len, d_model, n_heads = 2, 10, 64, 8
# 1. 测试注意力
q = k = v = torch.randn(batch, n_heads, seq_len, d_model // n_heads)
out, attn = attention(q, k, v)
assert out.shape == (batch, n_heads, seq_len, d_model // n_heads), "注意力输出形状错误"
assert attn.shape == (batch, n_heads, seq_len, seq_len), "注意力权重形状错误"
print("[OK] Scaled Dot-Product Attention")
# 2. 测试多头注意力
mha = MultiHeadAttention(n_heads, d_model)
x = torch.randn(batch, seq_len, d_model)
out = mha(x, x, x)
assert out.shape == (batch, seq_len, d_model), "多头注意力输出形状错误"
print("[OK] Multi-Head Attention")
# 3. 测试位置编码
pe = PositionalEncoding(d_model, 0.0)
out = pe(x)
assert out.shape == (batch, seq_len, d_model), "位置编码输出形状错误"
print("[OK] Positional Encoding")
# 4. 测试前馈网络
ff = PositionwiseFeedForward(d_model, d_model * 4)
out = ff(x)
assert out.shape == (batch, seq_len, d_model), "前馈网络输出形状错误"
print("[OK] Feed-Forward Network")
# 5. 测试 LayerNorm
ln = LayerNorm(d_model)
out = ln(x)
assert out.shape == (batch, seq_len, d_model), "LayerNorm 输出形状错误"
print("[OK] LayerNorm")
# 6. 测试因果掩码
mask = subsequent_mask(5)
assert mask.shape == (1, 5, 5), "因果掩码形状错误"
# 验证上三角被屏蔽
assert mask[0, 0, 1] == False, "因果掩码: 位置0不应看到位置1"
assert mask[0, 1, 0] == True, "因果掩码: 位置1应能看到位置0"
print("[OK] Subsequent Mask")
# 7. 测试完整模型
model = make_model(
src_vocab=100, tgt_vocab=100,
n_layers=2, d_model=64, d_ff=256, n_heads=4,
)
src = torch.randint(0, 100, (batch, seq_len))
tgt = torch.randint(0, 100, (batch, seq_len))
src_mask = torch.ones(batch, 1, seq_len)
tgt_mask = subsequent_mask(seq_len)
logits = model(src, tgt, src_mask, tgt_mask)
assert logits.shape == (batch, seq_len, 100), f"模型输出形状错误: {logits.shape}"
print("[OK] Full Transformer Model")
# 8. 测试梯度可正常反向传播
loss = logits.sum()
loss.backward()
grad_ok = all(p.grad is not None for p in model.parameters() if p.requires_grad)
assert grad_ok, "梯度反向传播失败"
print("[OK] Backpropagation")
print("=" * 60)
print("所有组件测试通过!")
print("=" * 60)
# ============================================================================
# 入口
# ============================================================================
if __name__ == "__main__":
# 1. 运行组件测试,验证各模块正确性
test_components()
# 2. 训练一个简单的字符级语言模型
train_language_model()
所有测试通过,训练和文本生成效果非常理想!总结如下:
完成情况
已在 cucode/decoder_only.py 中创建了一个 GPT 风格的精简版 Decoder-Only Transformer(类似 ChatGPT 的架构),共约 500 行代码,全部测试和训练验证通过。
与 trans.py 的关键区别
| 特性 | 完整 Transformer (trans.py) | Decoder-Only (本文件) |
|---|---|---|
| 编码器 | 有(双向自注意力) | 无 |
| 交叉注意力 | 有 | 无(只有因果自注意力) |
| 注意力方向 | 编码器双向/解码器单向 | 单向(只能看到过去) |
| 归一化位置 | Post-LN | Pre-LN(训练更稳定) |
| 激活函数 | ReLU | GELU |
| 位置编码 | 正弦(固定) | 可学习 |
文件结构
| 模块 | 说明 |
|---|---|
GPTConfig |
集中管理超参数(词表、上下文长度、层数、头数、维度) |
CausalSelfAttention |
因果自注意力:一次投影生成 QKV + 下三角掩码 |
MLP |
前馈网络(GELU 激活,4 倍升维) |
Block |
Transformer 块(Pre-LN + 残差连接) |
GPT |
完整模型:嵌入 → N 个 Block → LayerNorm → 输出层 |
make_gpt() |
模型构建函数 |
CharDataset / train_char_gpt() |
字符级训练示例 |
generate() |
自回归采样(支持温度、Top-K) |
test_components() |
8 项单元测试 |
GPT 特有的工程技巧(学习重点)
- 权重共享:
token_embedding.weight = lm_head.weight,输出层复用词嵌入矩阵,大幅减少参数量 - GPT-2 初始化:残差路径上
c_proj用0.02/√(2·n_layer)的小标准差,防止多层堆叠数值爆炸 - 因果掩码验证:专门测试"修改未来 token 不影响过去输出"的因果性质
- 生成技巧:温度缩放 + Top-K 采样,控制生成随机性
运行结果
组件测试: 全部 8 项通过
训练: Loss 2.2429 → 0.0908(下降 96%)
生成示例:
'hello ' -> 'hello world, this is a decoder only transforme'
'the quick' -> 'the quick brown fox jumps over the lazy dog. gpt '
'deep ' -> 'deep learning is fun, let us build a small gp'
模型仅 104,000 个参数(约 10 万),在 CPU 上 30 个 epoch 约 90 秒即学会了训练文本的模式——这正体现了 Decoder-Only 架构的本质:预测下一个 token。
运行方式:python cucode/decoder_only.py
"""
decoder_only.py
================
精简版 Decoder-Only Transformer(GPT 风格)从零实现(学习用途)。
本文件实现了 GPT-2 / ChatGPT 系列使用的核心架构,与完整 Transformer
(见 trans.py)的关键区别在于:
| 特性 | 完整 Transformer (trans.py) | Decoder-Only (本文件) |
|-----------------|-----------------------------|----------------------------|
| 编码器 | 有(双向自注意力) | 无 |
| 解码器 | 有(含交叉注意力) | 有(只有因果自注意力) |
| 注意力方向 | 编码器双向 / 解码器单向 | 只能看到过去(单向) |
| 归一化位置 | Post-LN(先加残差再归一化) | Pre-LN(先归一化再加残差) |
| 激活函数 | ReLU | GELU |
| 位置编码 | 正弦位置编码(固定) | 可学习位置编码 |
Decoder-Only 模型直接以"下一个 token 预测"为训练目标,因此天然适合
语言建模、对话生成等任务。ChatGPT / GPT 系列都采用这种架构。
运行环境: Python 3.10+, PyTorch 2.x
依赖: pip install torch
"""
import math
import time
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
# ============================================================================
# 第一部分:配置类
# ============================================================================
class GPTConfig:
"""GPT 模型的超参数配置(集中管理,方便修改)。"""
def __init__(
self,
vocab_size: int = 128, # 词表大小(token 种类数)
block_size: int = 64, # 最大上下文长度(能"看到"多少个 token)
n_layer: int = 2, # Transformer 块的数量
n_head: int = 4, # 注意力头数
n_embd: int = 128, # 嵌入维度(模型宽度)
dropout: float = 0.1, # dropout 比率
):
self.vocab_size = vocab_size
self.block_size = block_size
self.n_layer = n_layer
self.n_head = n_head
self.n_embd = n_embd
self.dropout = dropout
# ============================================================================
# 第二部分:核心组件
# ============================================================================
class CausalSelfAttention(nn.Module):
"""
因果自注意力 (Causal Self-Attention)。
与完整 Transformer 的多头注意力不同,它有两个特点:
1. 只有一个注意力层(对自身做注意力,Q=K=V 来自同一个输入);
2. 使用下三角掩码 (Causal Mask),使每个位置只能"看到"它自己
及之前的 token,防止看到未来信息(这是自回归生成的关键)。
实现技巧:Q、K、V 用一次线性投影同时生成(一次性矩阵乘法,效率更高),
输出时再用一次线性投影融合多头结果。
"""
def __init__(self, config: GPTConfig):
super().__init__()
assert config.n_embd % config.n_head == 0, "n_embd 必须能被 n_head 整除"
self.n_head = config.n_head
self.n_embd = config.n_embd
self.head_size = config.n_embd // config.n_head # 每个头的维度 d_k
# 一次线性投影同时生成 Q、K、V(输出维度是 n_embd 的 3 倍)
self.c_attn = nn.Linear(config.n_embd, 3 * config.n_embd)
# 输出投影
self.c_proj = nn.Linear(config.n_embd, config.n_embd)
self.attn_dropout = nn.Dropout(config.dropout)
self.resid_dropout = nn.Dropout(config.dropout)
# 因果掩码: 下三角矩阵 (block_size, block_size)
# tril 结果示例 (block_size=4):
# [[1, 0, 0, 0],
# [1, 1, 0, 0],
# [1, 1, 1, 0],
# [1, 1, 1, 1]]
# 注册为 buffer,不参与训练但会随模型移动设备
self.register_buffer(
"causal_mask",
torch.tril(torch.ones(config.block_size, config.block_size)).view(
1, 1, config.block_size, config.block_size
),
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""参数: x (batch, seq_len, n_embd),返回同形状。"""
B, T, C = x.size() # batch, 序列长度, 嵌入维度
# 1. 一次性生成 QKV 并切分
qkv = self.c_attn(x) # (B, T, 3C)
q, k, v = qkv.split(self.n_embd, dim=2) # 每个 (B, T, C)
# 2. 重塑为多头形式: (B, T, C) -> (B, n_head, T, head_size)
q = q.view(B, T, self.n_head, self.head_size).transpose(1, 2)
k = k.view(B, T, self.n_head, self.head_size).transpose(1, 2)
v = v.view(B, T, self.n_head, self.head_size).transpose(1, 2)
# 3. 缩放点积注意力: scores = Q @ K^T / sqrt(d_k)
# (B, n_head, T, head_size) @ (B, n_head, head_size, T) -> (B, n_head, T, T)
scores = q @ k.transpose(-2, -1) / math.sqrt(self.head_size)
# 4. 应用因果掩码: 未来位置设为 -inf,softmax 后权重趋近于 0
# 只取当前序列长度对应的掩码部分
scores = scores.masked_fill(
self.causal_mask[:, :, :T, :T] == 0, float("-inf")
)
# 5. softmax 归一化 + dropout 得到注意力权重,再对 V 加权求和
attn = F.softmax(scores, dim=-1)
attn = self.attn_dropout(attn)
y = attn @ v # (B, n_head, T, head_size)
# 6. 拼接所有头: (B, n_head, T, head_size) -> (B, T, C)
y = y.transpose(1, 2).contiguous().view(B, T, C)
# 7. 输出投影 + 残差 dropout
return self.resid_dropout(self.c_proj(y))
class MLP(nn.Module):
"""
前馈网络 (Feed-Forward Network),GPT 使用 GELU 激活。
MLP(x) = Linear(n_embd -> 4*n_embd) -> GELU -> Linear(4*n_embd -> n_embd)
为什么中间维度是 4 倍?这是 Transformer 论文中实验验证的经验值,
相当于给每个 token 一个"独立思考"的多层感知机。
"""
def __init__(self, config: GPTConfig):
super().__init__()
self.c_fc = nn.Linear(config.n_embd, 4 * config.n_embd) # 升维
self.c_proj = nn.Linear(4 * config.n_embd, config.n_embd) # 降维回 n_embd
self.gelu = nn.GELU() # GELU: 更平滑的 ReLU 变体
self.dropout = nn.Dropout(config.dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.c_fc(x)
x = self.gelu(x)
x = self.c_proj(x)
return self.dropout(x)
class Block(nn.Module):
"""
一个 Transformer 块:因果自注意力 + 前馈网络。
使用 Pre-LN(先 LayerNorm 再进子层,最后加残差),这是 GPT 系列的标准做法,
相比原论文的 Post-LN 更容易稳定训练:
x = x + Attn(LayerNorm(x)) # 子层1
x = x + MLP(LayerNorm(x)) # 子层2
"""
def __init__(self, config: GPTConfig):
super().__init__()
self.ln_1 = nn.LayerNorm(config.n_embd) # 注意力前的归一化
self.attn = CausalSelfAttention(config)
self.ln_2 = nn.LayerNorm(config.n_embd) # 前馈前的归一化
self.mlp = MLP(config)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = x + self.attn(self.ln_1(x)) # 残差 + 自注意力
x = x + self.mlp(self.ln_2(x)) # 残差 + 前馈网络
return x
# ============================================================================
# 第三部分:完整 GPT 模型
# ============================================================================
class GPT(nn.Module):
"""
Decoder-Only Transformer 模型(GPT 风格)。
数据流:
token_ids --[词嵌入]--> x
positions --[位置嵌入]--> pos
x = x + pos
for each Block: x = Block(x) # 堆叠 n_layer 个 Transformer 块
x = LayerNorm(x)
logits = Linear(x) # 映射到词表概率
"""
def __init__(self, config: GPTConfig):
super().__init__()
self.config = config
# 1. 词嵌入: token id -> n_embd 维向量
self.token_embedding = nn.Embedding(config.vocab_size, config.n_embd)
# 2. 位置嵌入: 位置索引 -> n_embd 维向量(可学习的)
self.position_embedding = nn.Embedding(config.block_size, config.n_embd)
# 3. 堆叠的 Transformer 块
self.blocks = nn.ModuleList([Block(config) for _ in range(config.n_layer)])
# 4. 最终 LayerNorm
self.ln_f = nn.LayerNorm(config.n_embd)
# 5. 输出层: n_embd -> vocab_size
self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)
# 技巧: 权重共享 —— 让输出层复用词嵌入的权重矩阵
# 因为"预测下一个词"和"查词向量"本质是同一张表,共享可大幅减少参数量
self.token_embedding.weight = self.lm_head.weight
# 参数初始化
self.apply(self._init_weights)
# GPT-2 的技巧: 对残差路径上的线性层用更小的标准差初始化
# 防止多块堆叠后数值过大
for name, p in self.named_parameters():
if name.endswith("c_proj.weight"):
nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * config.n_layer))
def _init_weights(self, module: nn.Module):
"""初始化权重: 线性层/嵌入层用 N(0, 0.02),偏置与归一化层置零。"""
if isinstance(module, (nn.Linear, nn.Embedding)):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
if isinstance(module, nn.Linear) and module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.LayerNorm):
nn.init.zeros_(module.bias)
nn.init.ones_(module.weight)
def forward(
self,
idx: torch.Tensor,
targets: torch.Tensor | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
"""
参数:
idx: 输入 token id 序列 (batch, seq_len)
targets: 目标 token id 序列 (batch, seq_len),训练时提供
用于计算交叉熵损失;推理时为 None
返回:
训练时: (logits, loss)
推理时: (logits,)
"""
B, T = idx.size()
assert T <= self.config.block_size, f"序列长度 {T} 超过上下文上限 {self.config.block_size}"
# 1. 词嵌入 (B, T, n_embd)
tok_emb = self.token_embedding(idx)
# 2. 位置嵌入 (T, n_embd)
pos = torch.arange(0, T, dtype=torch.long, device=idx.device)
pos_emb = self.position_embedding(pos)
# 3. 相加得到输入表示
x = tok_emb + pos_emb
# 4. 经过所有 Transformer 块
for block in self.blocks:
x = block(x)
# 5. 最终归一化 + 投影到词表
x = self.ln_f(x)
logits = self.lm_head(x) # (B, T, vocab_size)
# 6. 训练时计算损失: 目标是"下一个 token",所以 logits 在位置 t
# 预测的是位置 t+1 的 token,与 targets 逐位比较
loss = None
if targets is not None:
loss = F.cross_entropy(
logits.view(-1, logits.size(-1)),
targets.view(-1),
ignore_index=-1,
)
return logits, loss
@torch.no_grad()
def generate(
self,
idx: torch.Tensor,
max_new_tokens: int = 50,
temperature: float = 1.0,
top_k: int | None = None,
) -> torch.Tensor:
"""
自回归文本生成。
逐 token 生成: 每次把已生成的序列喂回模型,只取最后一个位置的预测,
采样一个 token 追加到序列末尾,重复直到生成 max_new_tokens 个。
参数:
idx: 初始 prompt 的 token id (batch, seq_len)
max_new_tokens: 要生成的 token 数量
temperature: 采样温度。>1 更随机,<1 更确定,=0 取 argmax
top_k: 只从前 k 个概率最高的 token 中采样(可选)
返回:
完整序列 (batch, seq_len + max_new_tokens)
"""
for _ in range(max_new_tokens):
# 只取最后 block_size 个 token(防止超过上下文上限)
idx_cond = idx[:, -self.config.block_size:]
# 前向传播(推理模式,无需梯度)
logits, _ = self(idx_cond) # (B, T, vocab)
logits = logits[:, -1, :] # 只取最后一个位置 (B, vocab)
# 温度缩放
if temperature != 1.0:
logits = logits / temperature
# Top-K 过滤: 只保留概率最高的前 k 个
if top_k is not None:
v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
logits[logits < v[:, [-1]]] = float("-inf")
# 从概率分布中采样
probs = F.softmax(logits, dim=-1) # (B, vocab)
idx_next = torch.multinomial(probs, num_samples=1) # (B, 1)
# 追加到序列末尾
idx = torch.cat((idx, idx_next), dim=1)
return idx
# ============================================================================
# 第四部分:模型构建函数
# ============================================================================
def make_gpt(
vocab_size: int,
block_size: int = 64,
n_layer: int = 2,
n_head: int = 4,
n_embd: int = 128,
dropout: float = 0.1,
) -> GPT:
"""构建一个 GPT 模型(配置集中管理)。"""
config = GPTConfig(
vocab_size=vocab_size,
block_size=block_size,
n_layer=n_layer,
n_head=n_head,
n_embd=n_embd,
dropout=dropout,
)
model = GPT(config)
print(
f"模型构建完成 | 参数量: {sum(p.numel() for p in model.parameters()):,} "
f"| 层数: {n_layer} | 头数: {n_head} | 维度: {n_embd}"
)
return model
# ============================================================================
# 第五部分:字符级语言模型训练示例
# ============================================================================
class CharDataset(Dataset):
"""字符级数据集: 从文本中切分 (input, target) 训练样本。
GPT 的训练样本是"给定前 k 个字符,预测第 k+1 个字符"。
原始文本按块切分后,同一块内任意位置都天然构成训练样本
(因为因果注意力会让每个位置只看到自己之前的内容),
这里简单地对齐逐位作为 target。
"""
def __init__(self, text: str, block_size: int = 64):
self.chars = sorted(set(text))
self.char2idx = {c: i for i, c in enumerate(self.chars)}
self.idx2char = {i: c for c, i in self.char2idx.items()}
self.vocab_size = len(self.chars)
self.block_size = block_size
# 编码整个文本
self.data = torch.tensor([self.char2idx[c] for c in text], dtype=torch.long)
def __len__(self) -> int:
return max(0, len(self.data) - self.block_size)
def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor]:
# x: 连续的 block_size 个字符
x = self.data[idx : idx + self.block_size]
# y: 右移一位,即每个位置的"下一个字符"
y = self.data[idx + 1 : idx + self.block_size + 1]
return x, y
def train_char_gpt(seed: int = 42):
"""
用小型文本训练一个字符级 GPT 语言模型,演示完整训练流程。
训练目标是: 给定前面的字符序列,预测下一个字符。
这就是 ChatGPT 这类大模型的本质 —— 只是规模更大、数据更多。
"""
torch.manual_seed(seed)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"训练设备: {device}")
# --- 1. 准备数据 ---
sample_text = (
"hello world, this is a decoder only transformer. "
"it learns to predict the next character. "
"the quick brown fox jumps over the lazy dog. "
"gpt models are decoder only, just like chatgpt. "
"deep learning is fun, let us build a small gpt. "
) * 30 # 重复以增加数据量
block_size = 32 # 上下文长度(能看到的字符数)
dataset = CharDataset(sample_text, block_size=block_size)
dataloader = DataLoader(dataset, batch_size=32, shuffle=True)
print(f"词表大小: {dataset.vocab_size}, 字符: {''.join(dataset.chars)}")
print(f"训练样本数: {len(dataset)}")
# --- 2. 构建模型(小规模,便于 CPU 快速训练)---
model = make_gpt(
vocab_size=dataset.vocab_size,
block_size=block_size,
n_layer=2, # 2 层 Transformer 块
n_head=4, # 4 个注意力头
n_embd=64, # 嵌入维度 64
dropout=0.1,
).to(device)
# --- 3. 优化器 ---
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)
# --- 4. 训练循环 ---
epochs = 30
model.train()
print("\n开始训练...")
print("=" * 60)
for epoch in range(epochs):
t0 = time.time()
total_loss = 0.0
n_batches = 0
for x, y in dataloader:
x, y = x.to(device), y.to(device)
# 前向传播 + 损失
logits, loss = model(x, y)
# 反向传播
optimizer.zero_grad()
loss.backward()
optimizer.step()
total_loss += loss.item()
n_batches += 1
avg_loss = total_loss / max(n_batches, 1)
if (epoch + 1) % 5 == 0 or epoch == 0:
print(f"Epoch {epoch + 1:3d}/{epochs} | Loss: {avg_loss:.4f} | Time: {time.time() - t0:.2f}s")
print("=" * 60)
print("训练完成!")
# --- 5. 文本生成(自回归采样)---
print("\n--- 文本生成示例 ---")
model.eval()
prompts = ["hello ", "the quick", "deep "]
for prompt in prompts:
# 编码 prompt
ids = [dataset.char2idx.get(c, 0) for c in prompt]
idx = torch.tensor([ids], dtype=torch.long, device=device)
# 生成 40 个新字符(温度为 0.8,稍保守)
generated = model.generate(idx, max_new_tokens=40, temperature=0.8)
# 解码为文本
text = "".join(dataset.idx2char.get(i, "?") for i in generated[0].tolist())
print(f"提示: {prompt!r} -> {text!r}")
return model, dataset
# ============================================================================
# 第六部分:单元测试
# ============================================================================
def test_components():
"""运行断言测试,验证模型各组件的正确性。"""
print("\n" + "=" * 60)
print("运行组件测试...")
print("=" * 60)
batch, block_size = 2, 16
config = GPTConfig(
vocab_size=50, block_size=block_size,
n_layer=1, n_head=2, n_embd=32, dropout=0.0,
)
# 1. 测试因果自注意力
attn = CausalSelfAttention(config)
x = torch.randn(batch, block_size, 32)
out = attn(x)
assert out.shape == x.shape, f"注意力输出形状错误: {out.shape} != {x.shape}"
print("[OK] Causal Self-Attention")
# 2. 测试 MLP
mlp = MLP(config)
out = mlp(x)
assert out.shape == x.shape, f"MLP 输出形状错误: {out.shape} != {x.shape}"
print("[OK] MLP")
# 3. 测试 Block
block = Block(config)
out = block(x)
assert out.shape == x.shape, f"Block 输出形状错误: {out.shape} != {x.shape}"
print("[OK] Transformer Block")
# 4. 测试因果性(关键验证!)
# 检查第 t 个位置的输出是否不受未来位置影响:
# 改变输入最后一个 token,前 15 个位置的输出应该完全不变
attn.eval()
x1 = torch.randn(1, block_size, 32)
x2 = x1.clone()
x2[0, -1, :] = 999.0 # 大幅修改最后一个 token
with torch.no_grad():
o1 = attn(x1)
o2 = attn(x2)
assert torch.allclose(o1[0, :-1], o2[0, :-1], atol=1e-5), "因果性被破坏: 未来影响了过去"
print("[OK] Causal Property (未来不会影响过去)")
# 5. 测试完整 GPT 模型
gpt = GPT(config)
idx = torch.randint(0, 50, (batch, block_size))
logits, loss = gpt(idx, idx)
assert logits.shape == (batch, block_size, 50), f"模型输出形状错误: {logits.shape}"
assert loss is not None and loss.ndim == 0, "损失计算错误"
print("[OK] GPT Model Forward + Loss")
# 6. 测试推理模式(无 targets)
logits, loss = gpt(idx)
assert loss is None, "推理模式下 loss 应为 None"
print("[OK] GPT Inference Mode (loss=None)")
# 7. 测试梯度反向传播
loss = logits.sum()
loss.backward()
grad_ok = all(p.grad is not None for p in gpt.parameters() if p.requires_grad)
assert grad_ok, "梯度反向传播失败"
print("[OK] Backpropagation")
# 8. 测试生成
with torch.no_grad():
gen = gpt.generate(idx[:1], max_new_tokens=5, temperature=0.8)
assert gen.shape == (1, block_size + 5), f"生成序列形状错误: {gen.shape}"
print("[OK] Text Generation")
print("=" * 60)
print("所有组件测试通过!")
print("=" * 60)
# ============================================================================
# 入口
# ============================================================================
if __name__ == "__main__":
# 1. 运行组件测试
test_components()
# 2. 训练字符级 GPT 语言模型
train_char_gpt()
全部测试通过,训练和生成效果与 decoder_only.py 完全一致甚至更好!总结如下:
实际工程上用融合版MultiHead
"""
multi_head_atten_decode_only.py
================================
显式 Multi-Head Attention 版 Decoder-Only Transformer(GPT 风格)(学习用途)。
本文件实现与 decoder_only.py 完全相同的功能(GPT 架构 + 训练 + 生成),
唯一的区别在于**如何实现多头注意力**:
| 实现方式 | decoder_only.py | 本文件 |
|-----------------|----------------------------|-------------------------------|
| QKV 生成 | 一次线性投影生成全部 QKV | 每个头独立的 Q/K/V 线性层 |
| 多头处理 | reshape + transpose 切分 | 显式的 n_head 个 AttentionHead|
| 输出融合 | transpose + view 拼接 | torch.cat 拼接 + W_o 投影 |
| 教学价值 | 工程优化,GPU 友好 | 结构清晰,逐步可见 |
两种实现数学上等价(都是"缩放点积注意力 + 多头并行"),本文件用
"每个头一个独立模块"的方式把多头的概念显式展示出来:
MultiHeadAttention
├── head_0 (AttentionHead): w_q, w_k, w_v + scaled dot-product attention
├── head_1 (AttentionHead): w_q, w_k, w_v + scaled dot-product attention
├── ... (每个头负责关注不同方面的特征)
└── w_o: 拼接所有头的输出后做最终线性投影
补充: 工程实践中通常使用"融合版"实现(性能更高,见 FusedMultiHeadAttention):
FusedMultiHeadAttention
├── c_attn: nn.Linear(n_embd, 3*n_embd) # 一个大矩阵一次生成所有头的 QKV
├── view 拆分: (B, T, C) -> (B, n_head, T, head_size)
├── 批量注意力: (B, n_head, T, T) 所有头同时计算
└── c_proj: view 拼接后输出投影
为什么可以合并?所有头的 Q 投影矩阵按行拼接成一个大矩阵:
W_q_big = [W_q^1; W_q^2; ...; W_q^H] (d_model x d_model)
大矩阵乘法 x @ W_q_big^T 一次算出所有头的 Q,再 view 拆分,
结果与逐头独立计算完全一致(分块矩阵乘法)。
FLOPs 完全相同,但融合版实际运行更快,原因:
1. kernel launch 次数更少(H 次 GEMM -> 1 次 GEMM)
2. 大矩阵乘法更能利用 GPU 的 SIMD / Tensor Core
3. 中间张量更少,内存分配开销更低
运行环境: Python 3.10+, PyTorch 2.x
依赖: pip install torch
"""
import math
import time
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
# ============================================================================
# 第一部分:配置类
# ============================================================================
class GPTConfig:
"""GPT 模型的超参数配置(集中管理,方便修改)。"""
def __init__(
self,
vocab_size: int = 128, # 词表大小(token 种类数)
block_size: int = 64, # 最大上下文长度(能"看到"多少个 token)
n_layer: int = 2, # Transformer 块的数量
n_head: int = 4, # 注意力头数
n_embd: int = 128, # 嵌入维度(模型宽度)
dropout: float = 0.1, # dropout 比率
):
self.vocab_size = vocab_size
self.block_size = block_size
self.n_layer = n_layer
self.n_head = n_head
self.n_embd = n_embd
self.dropout = dropout
# ============================================================================
# 第二部分:显式 Multi-Head Attention
# ============================================================================
# 这是本文件与 decoder_only.py 的核心区别。
# 我们把"多头注意力"拆成三个层次,层层递进,方便理解:
#
# 层次 1: scaled_dot_product_attention() —— 单头注意力的核心计算
# 层次 2: AttentionHead —— 一个"头"(自己的 Q/K/V 投影)
# 层次 3: MultiHeadAttention —— 并行多个头 + 拼接 + 输出投影
def scaled_dot_product_attention(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
causal_mask: torch.Tensor,
dropout: nn.Dropout | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
缩放点积注意力 (Scaled Dot-Product Attention)——单个头的核心计算。
公式: Attention(Q, K, V) = softmax(Q @ K^T / sqrt(d_k)) @ V
参数:
q, k, v: (batch, seq_len, head_size)
causal_mask: (1, 1, block_size, block_size) 下三角掩码,
防止当前 token 看到未来的 token
dropout: 可选的注意力权重 dropout
返回:
output: 加权求和后的值 (batch, seq_len, head_size)
attn: 注意力权重矩阵 (batch, seq_len, seq_len)
"""
B, T, _ = q.shape
head_size = q.size(-1)
# 1. 计算注意力分数: Q @ K^T / sqrt(d_k)
# (B, T, head_size) @ (B, head_size, T) -> (B, T, T)
scores = q @ k.transpose(-2, -1) / math.sqrt(head_size)
# 2. 应用因果掩码: 未来位置设为 -inf,softmax 后权重趋近于 0
# 取当前序列长度对应的 (T x T) 部分,与 (B, T, T) 的 scores 广播
scores = scores.masked_fill(
causal_mask[0, 0, :T, :T] == 0, float("-inf")
)
# 3. softmax 归一化得到注意力权重
attn = F.softmax(scores, dim=-1)
# 4. 可选 dropout
if dropout is not None:
attn = dropout(attn)
# 5. 用注意力权重对 V 加权求和
out = attn @ v # (B, T, head_size)
return out, attn
class AttentionHead(nn.Module):
"""
单头注意力 (Attention Head)。
每个头拥有自己独立的 Q、K、V 线性投影矩阵,可以从不同角度
(不同子空间)观察输入序列。例如一个头关注"词性",另一个头
关注"指代关系",各司其职。
参数:
n_embd: 输入维度(模型宽度)
head_size: 该头的输出维度 d_k = n_embd / n_head
causal_mask: 因果掩码(下三角矩阵)
dropout: dropout 比率
"""
def __init__(self, n_embd: int, head_size: int, causal_mask: torch.Tensor, dropout: float = 0.1):
super().__init__()
self.head_size = head_size
self.causal_mask = causal_mask
# 每个头独立的 Q、K、V 线性投影
self.w_q = nn.Linear(n_embd, head_size, bias=False)
self.w_k = nn.Linear(n_embd, head_size, bias=False)
self.w_v = nn.Linear(n_embd, head_size, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
参数: x (batch, seq_len, n_embd)
返回: (batch, seq_len, head_size)
"""
# 1. 独立的 Q/K/V 投影(每个头用自己的权重矩阵)
q = self.w_q(x) # (B, T, head_size)
k = self.w_k(x)
v = self.w_v(x)
# 2. 缩放点积注意力
out, _ = scaled_dot_product_attention(q, k, v, self.causal_mask, self.dropout)
return out # (B, T, head_size)
class MultiHeadAttention(nn.Module):
"""
多头注意力 (Multi-Head Attention)——显式实现。
将输入同时送入 n_head 个并行的 AttentionHead,每个头在自己的子空间
独立做注意力,然后把所有头的输出**拼接**起来,最后通过输出投影 W_o
融合成一个 n_embd 维向量。
参数:
config: GPTConfig 配置
"""
def __init__(self, config: GPTConfig):
super().__init__()
assert config.n_embd % config.n_head == 0, "n_embd 必须能被 n_head 整除"
self.n_head = config.n_head
self.n_embd = config.n_embd
self.head_size = config.n_embd // config.n_head # 每个头的维度 d_k
# 因果掩码: 下三角矩阵 (1, 1, block_size, block_size)
# tril 结果示例 (block_size=4):
# [[1, 0, 0, 0],
# [1, 1, 0, 0],
# [1, 1, 1, 0],
# [1, 1, 1, 1]]
# 注册为 buffer,不参与训练但会随模型移动设备
self.register_buffer(
"causal_mask",
torch.tril(torch.ones(config.block_size, config.block_size)).view(
1, 1, config.block_size, config.block_size
),
)
# n_head 个并行的注意力头(各自拥有独立的 Q/K/V 权重)
self.heads = nn.ModuleList(
[
AttentionHead(config.n_embd, self.head_size, self.causal_mask, config.dropout)
for _ in range(config.n_head)
]
)
# 输出投影 W_o: 拼接后的 n_embd 维 -> n_embd 维
self.w_o = nn.Linear(config.n_embd, config.n_embd)
self.resid_dropout = nn.Dropout(config.dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
参数: x (batch, seq_len, n_embd)
返回: (batch, seq_len, n_embd)
"""
# 1. 每个头独立计算(并行执行,互不影响)
head_outputs = [head(x) for head in self.heads] # n_head 个 (B, T, head_size)
# 2. 沿特征维度拼接所有头的输出
# (B, T, head_size) x n_head -> (B, T, n_head * head_size) = (B, T, n_embd)
y = torch.cat(head_outputs, dim=-1)
# 3. 输出投影融合多头信息
return self.resid_dropout(self.w_o(y))
class FusedMultiHeadAttention(nn.Module):
"""
融合版多头注意力(工程优化实现,性能更高)。
与显式版 MultiHeadAttention 的区别:
| | 显式版 MultiHeadAttention | 融合版 FusedMultiHeadAttention |
|-----------------|---------------------------|--------------------------------|
| QKV 投影 | H 个独立的 w_q/w_k/w_v | 一个大线性层 c_attn 一次生成 |
| 多头拆分 | 每个头一个 AttentionHead | view + transpose 按块切分 |
| 注意力计算 | 每个头单独循环计算 | (B, n_head, T, T) 批量一次算完 |
| 输出融合 | torch.cat 拼接 + w_o | view 拼接 + c_proj |
数学等价的关键:
所有头的 Q 投影矩阵按行拼接 [W_q^1; ...; W_q^H] 得到一个大矩阵 W_q_big
(d_model x d_model)。大矩阵乘法 x @ W_q_big^T 一次算出全部头的 Q,
再 view 拆分,与逐头计算完全一致(分块矩阵乘法)。
性能优势(FLOPs 相同,但实际更快):
1. kernel launch 次数更少: H 次小 GEMM -> 1 次大 GEMM
2. 大矩阵乘法更能利用 GPU 的 SIMD / Tensor Core 算力
3. 中间张量更少,显存分配开销更低
本文件中两种实现都保留,并通过 test_fused_equivalence 验证其数学等价。
"""
def __init__(self, config: GPTConfig):
super().__init__()
assert config.n_embd % config.n_head == 0, "n_embd 必须能被 n_head 整除"
self.n_head = config.n_head
self.n_embd = config.n_embd
self.head_size = config.n_embd // config.n_head # 每个头的维度 d_k
# 一个大线性层同时生成 Q、K、V(输出维度是 n_embd 的 3 倍)
# 内部等价于把 H 个头各自的 [W_q^h; W_k^h; W_v^h] 按行拼接成一个大矩阵
# bias=False 与显式版各头(w_q/w_k/w_v 无 bias)保持参数完全一致
self.c_attn = nn.Linear(config.n_embd, 3 * config.n_embd, bias=False)
# 输出投影
self.c_proj = nn.Linear(config.n_embd, config.n_embd)
# 因果掩码(下三角矩阵)
self.register_buffer(
"causal_mask",
torch.tril(torch.ones(config.block_size, config.block_size)).view(
1, 1, config.block_size, config.block_size
),
)
self.attn_dropout = nn.Dropout(config.dropout)
self.resid_dropout = nn.Dropout(config.dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
参数: x (batch, seq_len, n_embd),返回同形状。
"""
B, T, C = x.size() # batch, 序列长度, 嵌入维度
# 1. 一次生成全部 QKV 并切分为 Q/K/V 三个部分
qkv = self.c_attn(x) # (B, T, 3C)
q, k, v = qkv.split(self.n_embd, dim=2) # 每个 (B, T, C)
# 2. view 拆分为多头形式(关键一步: 大矩阵结果的"按块切分")
# (B, T, C) -> (B, T, n_head, head_size) -> (B, n_head, T, head_size)
# head_size 维上第 h 个块就是第 h 个头的结果
q = q.view(B, T, self.n_head, self.head_size).transpose(1, 2)
k = k.view(B, T, self.n_head, self.head_size).transpose(1, 2)
v = v.view(B, T, self.n_head, self.head_size).transpose(1, 2)
# 3. 缩放点积注意力(所有头作为一个 batch 同时计算,无需循环!)
# (B, n_head, T, head_size) @ (B, n_head, head_size, T) -> (B, n_head, T, T)
scores = q @ k.transpose(-2, -1) / math.sqrt(self.head_size)
# 4. 应用因果掩码
scores = scores.masked_fill(
self.causal_mask[:, :, :T, :T] == 0, float("-inf")
)
# 5. softmax + dropout + 对 V 加权求和
attn = F.softmax(scores, dim=-1)
attn = self.attn_dropout(attn)
y = attn @ v # (B, n_head, T, head_size)
# 6. 拼接所有头(view 还原): (B, n_head, T, head_size) -> (B, T, C)
y = y.transpose(1, 2).contiguous().view(B, T, C)
# 7. 输出投影 + 残差 dropout
return self.resid_dropout(self.c_proj(y))
# ============================================================================
# 第三部分:前馈网络与 Transformer 块
# ============================================================================
class MLP(nn.Module):
"""
前馈网络 (Feed-Forward Network),GPT 使用 GELU 激活。
MLP(x) = Linear(n_embd -> 4*n_embd) -> GELU -> Linear(4*n_embd -> n_embd)
为什么中间维度是 4 倍?这是 Transformer 论文中实验验证的经验值,
相当于给每个 token 一个"独立思考"的多层感知机。
"""
def __init__(self, config: GPTConfig):
super().__init__()
self.c_fc = nn.Linear(config.n_embd, 4 * config.n_embd) # 升维
self.c_proj = nn.Linear(4 * config.n_embd, config.n_embd) # 降维回 n_embd
self.gelu = nn.GELU() # GELU: 更平滑的 ReLU 变体
self.dropout = nn.Dropout(config.dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.c_fc(x)
x = self.gelu(x)
x = self.c_proj(x)
return self.dropout(x)
class Block(nn.Module):
"""
一个 Transformer 块:多头注意力 + 前馈网络。
使用 Pre-LN(先 LayerNorm 再进子层,最后加残差),这是 GPT 系列的标准做法,
相比原论文的 Post-LN 更容易稳定训练:
x = x + MultiHeadAttention(LayerNorm(x)) # 子层1
x = x + MLP(LayerNorm(x)) # 子层2
参数:
config: GPTConfig 配置
use_fused: True 使用融合版注意力(FusedMultiHeadAttention,性能更高),
False 使用显式版(MultiHeadAttention,结构更清晰,默认)
"""
def __init__(self, config: GPTConfig, use_fused: bool = False):
super().__init__()
self.ln_1 = nn.LayerNorm(config.n_embd) # 注意力前的归一化
# 两种多头注意力实现任选其一(数学等价)
self.attn = FusedMultiHeadAttention(config) if use_fused else MultiHeadAttention(config)
self.ln_2 = nn.LayerNorm(config.n_embd) # 前馈前的归一化
self.mlp = MLP(config)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = x + self.attn(self.ln_1(x)) # 残差 + 多头自注意力
x = x + self.mlp(self.ln_2(x)) # 残差 + 前馈网络
return x
# ============================================================================
# 第四部分:完整 GPT 模型
# ============================================================================
class GPT(nn.Module):
"""
Decoder-Only Transformer 模型(GPT 风格),使用显式多头注意力。
数据流:
token_ids --[词嵌入]--> x
positions --[位置嵌入]--> pos
x = x + pos
for each Block: x = Block(x) # 堆叠 n_layer 个 Transformer 块
x = LayerNorm(x)
logits = Linear(x) # 映射到词表概率
"""
def __init__(self, config: GPTConfig, use_fused: bool = False):
super().__init__()
self.config = config
# 1. 词嵌入: token id -> n_embd 维向量
self.token_embedding = nn.Embedding(config.vocab_size, config.n_embd)
# 2. 位置嵌入: 位置索引 -> n_embd 维向量(可学习的)
self.position_embedding = nn.Embedding(config.block_size, config.n_embd)
# 3. 堆叠的 Transformer 块(use_fused=True 时用融合版注意力)
self.blocks = nn.ModuleList(
[Block(config, use_fused=use_fused) for _ in range(config.n_layer)]
)
# 4. 最终 LayerNorm
self.ln_f = nn.LayerNorm(config.n_embd)
# 5. 输出层: n_embd -> vocab_size
self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)
# 技巧: 权重共享 —— 让输出层复用词嵌入的权重矩阵
# 因为"预测下一个词"和"查词向量"本质是同一张表,共享可大幅减少参数量
self.token_embedding.weight = self.lm_head.weight
# 参数初始化
self.apply(self._init_weights)
# GPT-2 的技巧: 对残差路径上的线性层(输出投影)用更小的标准差初始化
# 防止多块堆叠后数值过大
for name, p in self.named_parameters():
if name.endswith("w_o.weight") or name.endswith("c_proj.weight"):
nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * config.n_layer))
def _init_weights(self, module: nn.Module):
"""初始化权重: 线性层/嵌入层用 N(0, 0.02),偏置与归一化层置零。"""
if isinstance(module, (nn.Linear, nn.Embedding)):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
if isinstance(module, nn.Linear) and module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.LayerNorm):
nn.init.zeros_(module.bias)
nn.init.ones_(module.weight)
def forward(
self,
idx: torch.Tensor,
targets: torch.Tensor | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
"""
参数:
idx: 输入 token id 序列 (batch, seq_len)
targets: 目标 token id 序列 (batch, seq_len),训练时提供
用于计算交叉熵损失;推理时为 None
返回:
训练时: (logits, loss)
推理时: (logits,)
"""
B, T = idx.size()
assert T <= self.config.block_size, f"序列长度 {T} 超过上下文上限 {self.config.block_size}"
# 1. 词嵌入 (B, T, n_embd)
tok_emb = self.token_embedding(idx)
# 2. 位置嵌入 (T, n_embd)
pos = torch.arange(0, T, dtype=torch.long, device=idx.device)
pos_emb = self.position_embedding(pos)
# 3. 相加得到输入表示
x = tok_emb + pos_emb
# 4. 经过所有 Transformer 块
for block in self.blocks:
x = block(x)
# 5. 最终归一化 + 投影到词表
x = self.ln_f(x)
logits = self.lm_head(x) # (B, T, vocab_size)
# 6. 训练时计算损失: 目标是"下一个 token",所以 logits 在位置 t
# 预测的是位置 t+1 的 token,与 targets 逐位比较
loss = None
if targets is not None:
loss = F.cross_entropy(
logits.view(-1, logits.size(-1)),
targets.view(-1),
ignore_index=-1,
)
return logits, loss
@torch.no_grad()
def generate(
self,
idx: torch.Tensor,
max_new_tokens: int = 50,
temperature: float = 1.0,
top_k: int | None = None,
) -> torch.Tensor:
"""
自回归文本生成。
逐 token 生成: 每次把已生成的序列喂回模型,只取最后一个位置的预测,
采样一个 token 追加到序列末尾,重复直到生成 max_new_tokens 个。
参数:
idx: 初始 prompt 的 token id (batch, seq_len)
max_new_tokens: 要生成的 token 数量
temperature: 采样温度。>1 更随机,<1 更确定,=0 取 argmax
top_k: 只从前 k 个概率最高的 token 中采样(可选)
返回:
完整序列 (batch, seq_len + max_new_tokens)
"""
for _ in range(max_new_tokens):
# 只取最后 block_size 个 token(防止超过上下文上限)
idx_cond = idx[:, -self.config.block_size:]
# 前向传播(推理模式,无需梯度)
logits, _ = self(idx_cond) # (B, T, vocab)
logits = logits[:, -1, :] # 只取最后一个位置 (B, vocab)
# 温度缩放
if temperature != 1.0:
logits = logits / temperature
# Top-K 过滤: 只保留概率最高的前 k 个
if top_k is not None:
v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
logits[logits < v[:, [-1]]] = float("-inf")
# 从概率分布中采样
probs = F.softmax(logits, dim=-1) # (B, vocab)
idx_next = torch.multinomial(probs, num_samples=1) # (B, 1)
# 追加到序列末尾
idx = torch.cat((idx, idx_next), dim=1)
return idx
# ============================================================================
# 第五部分:模型构建函数
# ============================================================================
def make_gpt(
vocab_size: int,
block_size: int = 64,
n_layer: int = 2,
n_head: int = 4,
n_embd: int = 128,
dropout: float = 0.1,
use_fused: bool = False,
) -> GPT:
"""构建一个 GPT 模型(配置集中管理)。
参数:
use_fused: True 使用融合版注意力(性能更高),False 使用显式版(默认)
"""
config = GPTConfig(
vocab_size=vocab_size,
block_size=block_size,
n_layer=n_layer,
n_head=n_head,
n_embd=n_embd,
dropout=dropout,
)
model = GPT(config, use_fused=use_fused)
print(
f"模型构建完成 | 参数量: {sum(p.numel() for p in model.parameters()):,} "
f"| 层数: {n_layer} | 头数: {n_head} | 维度: {n_embd} "
f"| 注意力: {'融合版' if use_fused else '显式版'}"
)
return model
# ============================================================================
# 第六部分:字符级语言模型训练示例
# ============================================================================
class CharDataset(Dataset):
"""字符级数据集: 从文本中切分 (input, target) 训练样本。
GPT 的训练样本是"给定前 k 个字符,预测第 k+1 个字符"。
原始文本按块切分后,同一块内任意位置都天然构成训练样本
(因为因果注意力会让每个位置只看到自己之前的内容),
这里简单地对齐逐位作为 target。
"""
def __init__(self, text: str, block_size: int = 64):
self.chars = sorted(set(text))
self.char2idx = {c: i for i, c in enumerate(self.chars)}
self.idx2char = {i: c for c, i in self.char2idx.items()}
self.vocab_size = len(self.chars)
self.block_size = block_size
# 编码整个文本
self.data = torch.tensor([self.char2idx[c] for c in text], dtype=torch.long)
def __len__(self) -> int:
return max(0, len(self.data) - self.block_size)
def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor]:
# x: 连续的 block_size 个字符
x = self.data[idx : idx + self.block_size]
# y: 右移一位,即每个位置的"下一个字符"
y = self.data[idx + 1 : idx + self.block_size + 1]
return x, y
def train_char_gpt(seed: int = 42):
"""
用小型文本训练一个字符级 GPT 语言模型,演示完整训练流程。
训练目标是: 给定前面的字符序列,预测下一个字符。
这就是 ChatGPT 这类大模型的本质 —— 只是规模更大、数据更多。
"""
torch.manual_seed(seed)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"训练设备: {device}")
# --- 1. 准备数据 ---
sample_text = (
"hello world, this is a decoder only transformer. "
"it learns to predict the next character. "
"the quick brown fox jumps over the lazy dog. "
"gpt models are decoder only, just like chatgpt. "
"deep learning is fun, let us build a small gpt. "
) * 30 # 重复以增加数据量
block_size = 32 # 上下文长度(能看到的字符数)
dataset = CharDataset(sample_text, block_size=block_size)
dataloader = DataLoader(dataset, batch_size=32, shuffle=True)
print(f"词表大小: {dataset.vocab_size}, 字符: {''.join(dataset.chars)}")
print(f"训练样本数: {len(dataset)}")
# --- 2. 构建模型(小规模,便于 CPU 快速训练)---
model = make_gpt(
vocab_size=dataset.vocab_size,
block_size=block_size,
n_layer=2, # 2 层 Transformer 块
n_head=4, # 4 个注意力头
n_embd=64, # 嵌入维度 64
dropout=0.1,
).to(device)
# --- 3. 优化器 ---
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)
# --- 4. 训练循环 ---
epochs = 30
model.train()
print("\n开始训练...")
print("=" * 60)
for epoch in range(epochs):
t0 = time.time()
total_loss = 0.0
n_batches = 0
for x, y in dataloader:
x, y = x.to(device), y.to(device)
# 前向传播 + 损失
logits, loss = model(x, y)
# 反向传播
optimizer.zero_grad()
loss.backward()
optimizer.step()
total_loss += loss.item()
n_batches += 1
avg_loss = total_loss / max(n_batches, 1)
if (epoch + 1) % 5 == 0 or epoch == 0:
print(f"Epoch {epoch + 1:3d}/{epochs} | Loss: {avg_loss:.4f} | Time: {time.time() - t0:.2f}s")
print("=" * 60)
print("训练完成!")
# --- 5. 文本生成(自回归采样)---
print("\n--- 文本生成示例 ---")
model.eval()
prompts = ["hello ", "the quick", "deep "]
for prompt in prompts:
# 编码 prompt
ids = [dataset.char2idx.get(c, 0) for c in prompt]
idx = torch.tensor([ids], dtype=torch.long, device=device)
# 生成 40 个新字符(温度为 0.8,稍保守)
generated = model.generate(idx, max_new_tokens=40, temperature=0.8)
# 解码为文本
text = "".join(dataset.idx2char.get(i, "?") for i in generated[0].tolist())
print(f"提示: {prompt!r} -> {text!r}")
return model, dataset
# ============================================================================
# 第七部分:单元测试
# ============================================================================
def test_components():
"""运行断言测试,验证模型各组件的正确性。"""
print("\n" + "=" * 60)
print("运行组件测试...")
print("=" * 60)
batch, block_size = 2, 16
config = GPTConfig(
vocab_size=50, block_size=block_size,
n_layer=1, n_head=2, n_embd=32, dropout=0.0,
)
# 1. 测试单头注意力(AttentionHead)
causal_mask = torch.tril(torch.ones(1, 1, block_size, block_size))
head = AttentionHead(n_embd=32, head_size=16, causal_mask=causal_mask)
x = torch.randn(batch, block_size, 32)
out = head(x)
assert out.shape == (batch, block_size, 16), f"单头注意力输出形状错误: {out.shape}"
print("[OK] AttentionHead (单头注意力)")
# 2. 测试多头注意力(MultiHeadAttention)
mha = MultiHeadAttention(config)
out = mha(x)
assert out.shape == x.shape, f"多头注意力输出形状错误: {out.shape} != {x.shape}"
print("[OK] MultiHeadAttention (多头注意力)")
# 2b. 测试融合版多头注意力(FusedMultiHeadAttention)
fused = FusedMultiHeadAttention(config)
out = fused(x)
assert out.shape == x.shape, f"融合版多头注意力输出形状错误: {out.shape} != {x.shape}"
print("[OK] FusedMultiHeadAttention (融合版多头注意力)")
# 2c. 融合版参数数量 = 显式版参数数量(数学等价性的数量体现)
n_fused = sum(p.numel() for p in fused.parameters())
n_explicit = sum(p.numel() for p in mha.parameters())
assert n_fused == n_explicit, f"两种实现的参数量不一致: {n_fused} != {n_explicit}"
print(f"[OK] Parameter Count Equal (融合版 {n_fused:,} == 显式版 {n_explicit:,})")
# 3. 测试多头拼接维度(核心验证!)
# 确认每个头输出 head_size 维,拼接后正好是 n_embd 维
head_outputs = [head(x) for head in mha.heads]
concat = torch.cat(head_outputs, dim=-1)
assert concat.shape[-1] == config.n_head * mha.head_size == config.n_embd, \
f"多头拼接维度错误: {concat.shape[-1]} != {config.n_embd}"
print(f"[OK] Multi-Head Concatenation ({config.n_head} 个头拼接 = {concat.shape[-1]} 维)")
# 4. 测试 MLP
mlp = MLP(config)
out = mlp(x)
assert out.shape == x.shape, f"MLP 输出形状错误: {out.shape} != {x.shape}"
print("[OK] MLP")
# 5. 测试 Block
block = Block(config)
out = block(x)
assert out.shape == x.shape, f"Block 输出形状错误: {out.shape} != {x.shape}"
print("[OK] Transformer Block")
# 6. 测试因果性(关键验证!)
# 检查第 t 个位置的输出是否不受未来位置影响:
# 改变输入最后一个 token,前 15 个位置的输出应该完全不变
mha.eval()
x1 = torch.randn(1, block_size, 32)
x2 = x1.clone()
x2[0, -1, :] = 999.0 # 大幅修改最后一个 token
with torch.no_grad():
o1 = mha(x1)
o2 = mha(x2)
assert torch.allclose(o1[0, :-1], o2[0, :-1], atol=1e-5), "因果性被破坏: 未来影响了过去"
print("[OK] Causal Property (未来不会影响过去)")
# 7. 测试完整 GPT 模型
gpt = GPT(config)
idx = torch.randint(0, 50, (batch, block_size))
logits, loss = gpt(idx, idx)
assert logits.shape == (batch, block_size, 50), f"模型输出形状错误: {logits.shape}"
assert loss is not None and loss.ndim == 0, "损失计算错误"
print("[OK] GPT Model Forward + Loss")
# 8. 测试推理模式(无 targets)
logits, loss = gpt(idx)
assert loss is None, "推理模式下 loss 应为 None"
print("[OK] GPT Inference Mode (loss=None)")
# 9. 测试梯度反向传播
loss = logits.sum()
loss.backward()
grad_ok = all(p.grad is not None for p in gpt.parameters() if p.requires_grad)
assert grad_ok, "梯度反向传播失败"
print("[OK] Backpropagation")
# 10. 测试生成
with torch.no_grad():
gen = gpt.generate(idx[:1], max_new_tokens=5, temperature=0.8)
assert gen.shape == (1, block_size + 5), f"生成序列形状错误: {gen.shape}"
print("[OK] Text Generation")
print("=" * 60)
print("所有组件测试通过!")
print("=" * 60)
def test_fused_equivalence():
"""
核心验证: 融合版(一个大线性层 + view 拆分)与显式版(多个独立头)
在权重对齐的前提下,前向输出是否完全一致。
验证方法(这也是"合并"的数学桥梁):
融合版 c_attn 的权重矩阵 (3C, C) 按行划分:
- 前 C 行 = 所有头的 Q 投影拼接成的大矩阵 W_q_big
- 中间 C 行 = 所有头的 K 投影拼接成的大矩阵 W_k_big
- 后 C 行 = 所有头的 V 投影拼接成的大矩阵 W_v_big
其中每个头的 Q 权重就是 W_q_big 中连续的 head_size 行:
head_h.w_q.weight = W_q_big[h*head_size : (h+1)*head_size]
把这些行赋给显式版的各个 AttentionHead,两者输出应完全一致。
"""
print("\n" + "=" * 60)
print("运行融合版 vs 显式版数学等价性验证...")
print("=" * 60)
config = GPTConfig(
vocab_size=50, block_size=16,
n_layer=1, n_head=4, n_embd=32, dropout=0.0,
)
C, H, hs = config.n_embd, config.n_head, config.n_embd // config.n_head
fused = FusedMultiHeadAttention(config)
explicit = MultiHeadAttention(config)
# 1. 对齐权重: 从融合版大矩阵中"按块切出"每个头的 Q/K/V 权重
with torch.no_grad():
w = fused.c_attn.weight # (3C, C),c_attn 无 bias
for h in range(H):
head = explicit.heads[h]
# 第 h 个头: 大矩阵中 [h*hs : (h+1)*hs] 行对应 Q/K/V 部分
head.w_q.weight.data.copy_(w[0 * C + h * hs : 0 * C + (h + 1) * hs])
head.w_k.weight.data.copy_(w[1 * C + h * hs : 1 * C + (h + 1) * hs])
head.w_v.weight.data.copy_(w[2 * C + h * hs : 2 * C + (h + 1) * hs])
# 对齐输出投影
explicit.w_o.weight.data.copy_(fused.c_proj.weight)
explicit.w_o.bias.data.copy_(fused.c_proj.bias)
# 2. 用相同输入比较输出
torch.manual_seed(0)
x = torch.randn(2, 16, C)
with torch.no_grad():
y_fused = fused(x)
y_explicit = explicit(x)
# 3. 断言完全一致
assert y_fused.shape == y_explicit.shape, "两种实现输出形状不一致"
diff = (y_fused - y_explicit).abs().max().item()
assert torch.allclose(y_fused, y_explicit, atol=1e-5), \
f"融合版与显式版输出不一致! 最大差异: {diff:.2e}"
print(f"[OK] Fused == Explicit | 输出最大差异: {diff:.2e} (阈值 1e-5)")
print(" 验证结论: 一个大线性层 + view 拆分 与 多个独立头 数学完全等价")
print("=" * 60)
# ============================================================================
# 入口
# ============================================================================
if __name__ == "__main__":
# 1. 运行组件测试
test_components()
# 2. 验证融合版与显式版数学等价
test_fused_equivalence()
# 3. 训练字符级 GPT 语言模型(默认显式版,可改 use_fused=True 体验融合版)
train_char_gpt()
完成情况
已在 cucode/multi_head_atten_decode_only.py 创建了显式 Multi-Head Attention 版的 Decoder-Only Transformer,功能与 decoder_only.py 完全一致,所有 10 项单元测试 + 训练 + 生成验证通过。
核心区别:显式的多头注意力实现
本文件把多头注意力拆成三个清晰层次,与 decoder_only.py 的"一次投影+reshape"优化实现形成对比:
MultiHeadAttention
├── head_0 (AttentionHead) ── w_q, w_k, w_v + 缩放点积注意力
├── head_1 (AttentionHead) ── w_q, w_k, w_v + 缩放点积注意力
├── head_2 (AttentionHead) ── ...
├── head_3 (AttentionHead) ── ...
└── w_o ── 拼接所有头的输出 → 输出投影
| 层次 | 类/函数 | 作用 |
|---|---|---|
| 1 | scaled_dot_product_attention() |
单头核心计算:softmax(QKᵀ/√dₖ)V |
| 2 | AttentionHead |
一个头:独立的 Q/K/V 线性投影 + 缩放点积注意力 |
| 3 | MultiHeadAttention |
并行 n 个 AttentionHead → torch.cat 拼接 → w_o 投影 |
与 decoder_only.py 的实现对比
| 实现方式 | decoder_only.py | 本文件 |
|---|---|---|
| QKV 生成 | 一次线性投影生成全部 QKV | 每个头独立的 w_q/w_k/w_v |
| 多头处理 | reshape + transpose 切分 | 显式的 nn.ModuleList 装 n 个头 |
| 输出融合 | transpose + view 拼接 | torch.cat + w_o 投影 |
| 教学价值 | 工程优化(GPU 高效) | 结构清晰、逐步可见 |
数学上两者完全等价,但本文件的实现让你能直观看到每个头在做什么——这正是理解 Transformer 核心思想的正确入门方式。
运行结果
组件测试: 全部 10 项通过(含因果性、多头拼接维度等关键验证)
模型参数量: 103,616 | Loss: 2.214 → 0.091
生成效果:
'hello ' -> 'hello world, this is a decoder only transforme'
'the quick' -> 'the quick brown fox jumps over the lazy dog. gpt '
'deep ' -> 'deep learning is fun, let us build a small gp'
建议学习路径
trans.py— 完整 Encoder-Decoder Transformer(seq2seq 架构)decoder_only.py— GPT 风格优化版多头注意力(工程写法)multi_head_atten_decode_only.py— 本文件,显式多头注意力(教学写法),建议先读这里再对比decoder_only.py,体会两种实现的等价关系
运行方式:python cucode/multi_head_atten_decode_only.py
"""
multi_head_atten_decode_only.py
================================
显式 Multi-Head Attention 版 Decoder-Only Transformer(GPT 风格)(学习用途)。
本文件实现与 decoder_only.py 完全相同的功能(GPT 架构 + 训练 + 生成),
唯一的区别在于**如何实现多头注意力**:
| 实现方式 | decoder_only.py | 本文件 |
|-----------------|----------------------------|-------------------------------|
| QKV 生成 | 一次线性投影生成全部 QKV | 每个头独立的 Q/K/V 线性层 |
| 多头处理 | reshape + transpose 切分 | 显式的 n_head 个 AttentionHead|
| 输出融合 | transpose + view 拼接 | torch.cat 拼接 + W_o 投影 |
| 教学价值 | 工程优化,GPU 友好 | 结构清晰,逐步可见 |
两种实现数学上等价(都是"缩放点积注意力 + 多头并行"),本文件用
"每个头一个独立模块"的方式把多头的概念显式展示出来:
MultiHeadAttention
├── head_0 (AttentionHead): w_q, w_k, w_v + scaled dot-product attention
├── head_1 (AttentionHead): w_q, w_k, w_v + scaled dot-product attention
├── ... (每个头负责关注不同方面的特征)
└── w_o: 拼接所有头的输出后做最终线性投影
运行环境: Python 3.10+, PyTorch 2.x
依赖: pip install torch
"""
import math
import time
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
# ============================================================================
# 第一部分:配置类
# ============================================================================
class GPTConfig:
"""GPT 模型的超参数配置(集中管理,方便修改)。"""
def __init__(
self,
vocab_size: int = 128, # 词表大小(token 种类数)
block_size: int = 64, # 最大上下文长度(能"看到"多少个 token)
n_layer: int = 2, # Transformer 块的数量
n_head: int = 4, # 注意力头数
n_embd: int = 128, # 嵌入维度(模型宽度)
dropout: float = 0.1, # dropout 比率
):
self.vocab_size = vocab_size
self.block_size = block_size
self.n_layer = n_layer
self.n_head = n_head
self.n_embd = n_embd
self.dropout = dropout
# ============================================================================
# 第二部分:显式 Multi-Head Attention
# ============================================================================
# 这是本文件与 decoder_only.py 的核心区别。
# 我们把"多头注意力"拆成三个层次,层层递进,方便理解:
#
# 层次 1: scaled_dot_product_attention() —— 单头注意力的核心计算
# 层次 2: AttentionHead —— 一个"头"(自己的 Q/K/V 投影)
# 层次 3: MultiHeadAttention —— 并行多个头 + 拼接 + 输出投影
def scaled_dot_product_attention(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
causal_mask: torch.Tensor,
dropout: nn.Dropout | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
缩放点积注意力 (Scaled Dot-Product Attention)——单个头的核心计算。
公式: Attention(Q, K, V) = softmax(Q @ K^T / sqrt(d_k)) @ V
参数:
q, k, v: (batch, seq_len, head_size)
causal_mask: (1, 1, block_size, block_size) 下三角掩码,
防止当前 token 看到未来的 token
dropout: 可选的注意力权重 dropout
返回:
output: 加权求和后的值 (batch, seq_len, head_size)
attn: 注意力权重矩阵 (batch, seq_len, seq_len)
"""
B, T, _ = q.shape
head_size = q.size(-1)
# 1. 计算注意力分数: Q @ K^T / sqrt(d_k)
# (B, T, head_size) @ (B, head_size, T) -> (B, T, T)
scores = q @ k.transpose(-2, -1) / math.sqrt(head_size)
# 2. 应用因果掩码: 未来位置设为 -inf,softmax 后权重趋近于 0
# 取当前序列长度对应的 (T x T) 部分,与 (B, T, T) 的 scores 广播
scores = scores.masked_fill(
causal_mask[0, 0, :T, :T] == 0, float("-inf")
)
# 3. softmax 归一化得到注意力权重
attn = F.softmax(scores, dim=-1)
# 4. 可选 dropout
if dropout is not None:
attn = dropout(attn)
# 5. 用注意力权重对 V 加权求和
out = attn @ v # (B, T, head_size)
return out, attn
class AttentionHead(nn.Module):
"""
单头注意力 (Attention Head)。
每个头拥有自己独立的 Q、K、V 线性投影矩阵,可以从不同角度
(不同子空间)观察输入序列。例如一个头关注"词性",另一个头
关注"指代关系",各司其职。
参数:
n_embd: 输入维度(模型宽度)
head_size: 该头的输出维度 d_k = n_embd / n_head
causal_mask: 因果掩码(下三角矩阵)
dropout: dropout 比率
"""
def __init__(self, n_embd: int, head_size: int, causal_mask: torch.Tensor, dropout: float = 0.1):
super().__init__()
self.head_size = head_size
self.causal_mask = causal_mask
# 每个头独立的 Q、K、V 线性投影
self.w_q = nn.Linear(n_embd, head_size, bias=False)
self.w_k = nn.Linear(n_embd, head_size, bias=False)
self.w_v = nn.Linear(n_embd, head_size, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
参数: x (batch, seq_len, n_embd)
返回: (batch, seq_len, head_size)
"""
# 1. 独立的 Q/K/V 投影(每个头用自己的权重矩阵)
q = self.w_q(x) # (B, T, head_size)
k = self.w_k(x)
v = self.w_v(x)
# 2. 缩放点积注意力
out, _ = scaled_dot_product_attention(q, k, v, self.causal_mask, self.dropout)
return out # (B, T, head_size)
class MultiHeadAttention(nn.Module):
"""
多头注意力 (Multi-Head Attention)——显式实现。
将输入同时送入 n_head 个并行的 AttentionHead,每个头在自己的子空间
独立做注意力,然后把所有头的输出**拼接**起来,最后通过输出投影 W_o
融合成一个 n_embd 维向量。
参数:
config: GPTConfig 配置
"""
def __init__(self, config: GPTConfig):
super().__init__()
assert config.n_embd % config.n_head == 0, "n_embd 必须能被 n_head 整除"
self.n_head = config.n_head
self.n_embd = config.n_embd
self.head_size = config.n_embd // config.n_head # 每个头的维度 d_k
# 因果掩码: 下三角矩阵 (1, 1, block_size, block_size)
# tril 结果示例 (block_size=4):
# [[1, 0, 0, 0],
# [1, 1, 0, 0],
# [1, 1, 1, 0],
# [1, 1, 1, 1]]
# 注册为 buffer,不参与训练但会随模型移动设备
self.register_buffer(
"causal_mask",
torch.tril(torch.ones(config.block_size, config.block_size)).view(
1, 1, config.block_size, config.block_size
),
)
# n_head 个并行的注意力头(各自拥有独立的 Q/K/V 权重)
self.heads = nn.ModuleList(
[
AttentionHead(config.n_embd, self.head_size, self.causal_mask, config.dropout)
for _ in range(config.n_head)
]
)
# 输出投影 W_o: 拼接后的 n_embd 维 -> n_embd 维
self.w_o = nn.Linear(config.n_embd, config.n_embd)
self.resid_dropout = nn.Dropout(config.dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
参数: x (batch, seq_len, n_embd)
返回: (batch, seq_len, n_embd)
"""
# 1. 每个头独立计算(并行执行,互不影响)
head_outputs = [head(x) for head in self.heads] # n_head 个 (B, T, head_size)
# 2. 沿特征维度拼接所有头的输出
# (B, T, head_size) x n_head -> (B, T, n_head * head_size) = (B, T, n_embd)
y = torch.cat(head_outputs, dim=-1)
# 3. 输出投影融合多头信息
return self.resid_dropout(self.w_o(y))
# ============================================================================
# 第三部分:前馈网络与 Transformer 块
# ============================================================================
class MLP(nn.Module):
"""
前馈网络 (Feed-Forward Network),GPT 使用 GELU 激活。
MLP(x) = Linear(n_embd -> 4*n_embd) -> GELU -> Linear(4*n_embd -> n_embd)
为什么中间维度是 4 倍?这是 Transformer 论文中实验验证的经验值,
相当于给每个 token 一个"独立思考"的多层感知机。
"""
def __init__(self, config: GPTConfig):
super().__init__()
self.c_fc = nn.Linear(config.n_embd, 4 * config.n_embd) # 升维
self.c_proj = nn.Linear(4 * config.n_embd, config.n_embd) # 降维回 n_embd
self.gelu = nn.GELU() # GELU: 更平滑的 ReLU 变体
self.dropout = nn.Dropout(config.dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.c_fc(x)
x = self.gelu(x)
x = self.c_proj(x)
return self.dropout(x)
class Block(nn.Module):
"""
一个 Transformer 块:显式多头注意力 + 前馈网络。
使用 Pre-LN(先 LayerNorm 再进子层,最后加残差),这是 GPT 系列的标准做法,
相比原论文的 Post-LN 更容易稳定训练:
x = x + MultiHeadAttention(LayerNorm(x)) # 子层1
x = x + MLP(LayerNorm(x)) # 子层2
"""
def __init__(self, config: GPTConfig):
super().__init__()
self.ln_1 = nn.LayerNorm(config.n_embd) # 注意力前的归一化
self.attn = MultiHeadAttention(config)
self.ln_2 = nn.LayerNorm(config.n_embd) # 前馈前的归一化
self.mlp = MLP(config)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = x + self.attn(self.ln_1(x)) # 残差 + 多头自注意力
x = x + self.mlp(self.ln_2(x)) # 残差 + 前馈网络
return x
# ============================================================================
# 第四部分:完整 GPT 模型
# ============================================================================
class GPT(nn.Module):
"""
Decoder-Only Transformer 模型(GPT 风格),使用显式多头注意力。
数据流:
token_ids --[词嵌入]--> x
positions --[位置嵌入]--> pos
x = x + pos
for each Block: x = Block(x) # 堆叠 n_layer 个 Transformer 块
x = LayerNorm(x)
logits = Linear(x) # 映射到词表概率
"""
def __init__(self, config: GPTConfig):
super().__init__()
self.config = config
# 1. 词嵌入: token id -> n_embd 维向量
self.token_embedding = nn.Embedding(config.vocab_size, config.n_embd)
# 2. 位置嵌入: 位置索引 -> n_embd 维向量(可学习的)
self.position_embedding = nn.Embedding(config.block_size, config.n_embd)
# 3. 堆叠的 Transformer 块
self.blocks = nn.ModuleList([Block(config) for _ in range(config.n_layer)])
# 4. 最终 LayerNorm
self.ln_f = nn.LayerNorm(config.n_embd)
# 5. 输出层: n_embd -> vocab_size
self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)
# 技巧: 权重共享 —— 让输出层复用词嵌入的权重矩阵
# 因为"预测下一个词"和"查词向量"本质是同一张表,共享可大幅减少参数量
self.token_embedding.weight = self.lm_head.weight
# 参数初始化
self.apply(self._init_weights)
# GPT-2 的技巧: 对残差路径上的线性层(输出投影)用更小的标准差初始化
# 防止多块堆叠后数值过大
for name, p in self.named_parameters():
if name.endswith("w_o.weight") or name.endswith("c_proj.weight"):
nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * config.n_layer))
def _init_weights(self, module: nn.Module):
"""初始化权重: 线性层/嵌入层用 N(0, 0.02),偏置与归一化层置零。"""
if isinstance(module, (nn.Linear, nn.Embedding)):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
if isinstance(module, nn.Linear) and module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.LayerNorm):
nn.init.zeros_(module.bias)
nn.init.ones_(module.weight)
def forward(
self,
idx: torch.Tensor,
targets: torch.Tensor | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
"""
参数:
idx: 输入 token id 序列 (batch, seq_len)
targets: 目标 token id 序列 (batch, seq_len),训练时提供
用于计算交叉熵损失;推理时为 None
返回:
训练时: (logits, loss)
推理时: (logits,)
"""
B, T = idx.size()
assert T <= self.config.block_size, f"序列长度 {T} 超过上下文上限 {self.config.block_size}"
# 1. 词嵌入 (B, T, n_embd)
tok_emb = self.token_embedding(idx)
# 2. 位置嵌入 (T, n_embd)
pos = torch.arange(0, T, dtype=torch.long, device=idx.device)
pos_emb = self.position_embedding(pos)
# 3. 相加得到输入表示
x = tok_emb + pos_emb
# 4. 经过所有 Transformer 块
for block in self.blocks:
x = block(x)
# 5. 最终归一化 + 投影到词表
x = self.ln_f(x)
logits = self.lm_head(x) # (B, T, vocab_size)
# 6. 训练时计算损失: 目标是"下一个 token",所以 logits 在位置 t
# 预测的是位置 t+1 的 token,与 targets 逐位比较
loss = None
if targets is not None:
loss = F.cross_entropy(
logits.view(-1, logits.size(-1)),
targets.view(-1),
ignore_index=-1,
)
return logits, loss
@torch.no_grad()
def generate(
self,
idx: torch.Tensor,
max_new_tokens: int = 50,
temperature: float = 1.0,
top_k: int | None = None,
) -> torch.Tensor:
"""
自回归文本生成。
逐 token 生成: 每次把已生成的序列喂回模型,只取最后一个位置的预测,
采样一个 token 追加到序列末尾,重复直到生成 max_new_tokens 个。
参数:
idx: 初始 prompt 的 token id (batch, seq_len)
max_new_tokens: 要生成的 token 数量
temperature: 采样温度。>1 更随机,<1 更确定,=0 取 argmax
top_k: 只从前 k 个概率最高的 token 中采样(可选)
返回:
完整序列 (batch, seq_len + max_new_tokens)
"""
for _ in range(max_new_tokens):
# 只取最后 block_size 个 token(防止超过上下文上限)
idx_cond = idx[:, -self.config.block_size:]
# 前向传播(推理模式,无需梯度)
logits, _ = self(idx_cond) # (B, T, vocab)
logits = logits[:, -1, :] # 只取最后一个位置 (B, vocab)
# 温度缩放
if temperature != 1.0:
logits = logits / temperature
# Top-K 过滤: 只保留概率最高的前 k 个
if top_k is not None:
v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
logits[logits < v[:, [-1]]] = float("-inf")
# 从概率分布中采样
probs = F.softmax(logits, dim=-1) # (B, vocab)
idx_next = torch.multinomial(probs, num_samples=1) # (B, 1)
# 追加到序列末尾
idx = torch.cat((idx, idx_next), dim=1)
return idx
# ============================================================================
# 第五部分:模型构建函数
# ============================================================================
def make_gpt(
vocab_size: int,
block_size: int = 64,
n_layer: int = 2,
n_head: int = 4,
n_embd: int = 128,
dropout: float = 0.1,
) -> GPT:
"""构建一个 GPT 模型(配置集中管理)。"""
config = GPTConfig(
vocab_size=vocab_size,
block_size=block_size,
n_layer=n_layer,
n_head=n_head,
n_embd=n_embd,
dropout=dropout,
)
model = GPT(config)
print(
f"模型构建完成 | 参数量: {sum(p.numel() for p in model.parameters()):,} "
f"| 层数: {n_layer} | 头数: {n_head} | 维度: {n_embd}"
)
return model
# ============================================================================
# 第六部分:字符级语言模型训练示例
# ============================================================================
class CharDataset(Dataset):
"""字符级数据集: 从文本中切分 (input, target) 训练样本。
GPT 的训练样本是"给定前 k 个字符,预测第 k+1 个字符"。
原始文本按块切分后,同一块内任意位置都天然构成训练样本
(因为因果注意力会让每个位置只看到自己之前的内容),
这里简单地对齐逐位作为 target。
"""
def __init__(self, text: str, block_size: int = 64):
self.chars = sorted(set(text))
self.char2idx = {c: i for i, c in enumerate(self.chars)}
self.idx2char = {i: c for c, i in self.char2idx.items()}
self.vocab_size = len(self.chars)
self.block_size = block_size
# 编码整个文本
self.data = torch.tensor([self.char2idx[c] for c in text], dtype=torch.long)
def __len__(self) -> int:
return max(0, len(self.data) - self.block_size)
def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor]:
# x: 连续的 block_size 个字符
x = self.data[idx : idx + self.block_size]
# y: 右移一位,即每个位置的"下一个字符"
y = self.data[idx + 1 : idx + self.block_size + 1]
return x, y
def train_char_gpt(seed: int = 42):
"""
用小型文本训练一个字符级 GPT 语言模型,演示完整训练流程。
训练目标是: 给定前面的字符序列,预测下一个字符。
这就是 ChatGPT 这类大模型的本质 —— 只是规模更大、数据更多。
"""
torch.manual_seed(seed)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"训练设备: {device}")
# --- 1. 准备数据 ---
sample_text = (
"hello world, this is a decoder only transformer. "
"it learns to predict the next character. "
"the quick brown fox jumps over the lazy dog. "
"gpt models are decoder only, just like chatgpt. "
"deep learning is fun, let us build a small gpt. "
) * 30 # 重复以增加数据量
block_size = 32 # 上下文长度(能看到的字符数)
dataset = CharDataset(sample_text, block_size=block_size)
dataloader = DataLoader(dataset, batch_size=32, shuffle=True)
print(f"词表大小: {dataset.vocab_size}, 字符: {''.join(dataset.chars)}")
print(f"训练样本数: {len(dataset)}")
# --- 2. 构建模型(小规模,便于 CPU 快速训练)---
model = make_gpt(
vocab_size=dataset.vocab_size,
block_size=block_size,
n_layer=2, # 2 层 Transformer 块
n_head=4, # 4 个注意力头
n_embd=64, # 嵌入维度 64
dropout=0.1,
).to(device)
# --- 3. 优化器 ---
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)
# --- 4. 训练循环 ---
epochs = 30
model.train()
print("\n开始训练...")
print("=" * 60)
for epoch in range(epochs):
t0 = time.time()
total_loss = 0.0
n_batches = 0
for x, y in dataloader:
x, y = x.to(device), y.to(device)
# 前向传播 + 损失
logits, loss = model(x, y)
# 反向传播
optimizer.zero_grad()
loss.backward()
optimizer.step()
total_loss += loss.item()
n_batches += 1
avg_loss = total_loss / max(n_batches, 1)
if (epoch + 1) % 5 == 0 or epoch == 0:
print(f"Epoch {epoch + 1:3d}/{epochs} | Loss: {avg_loss:.4f} | Time: {time.time() - t0:.2f}s")
print("=" * 60)
print("训练完成!")
# --- 5. 文本生成(自回归采样)---
print("\n--- 文本生成示例 ---")
model.eval()
prompts = ["hello ", "the quick", "deep "]
for prompt in prompts:
# 编码 prompt
ids = [dataset.char2idx.get(c, 0) for c in prompt]
idx = torch.tensor([ids], dtype=torch.long, device=device)
# 生成 40 个新字符(温度为 0.8,稍保守)
generated = model.generate(idx, max_new_tokens=40, temperature=0.8)
# 解码为文本
text = "".join(dataset.idx2char.get(i, "?") for i in generated[0].tolist())
print(f"提示: {prompt!r} -> {text!r}")
return model, dataset
# ============================================================================
# 第七部分:单元测试
# ============================================================================
def test_components():
"""运行断言测试,验证模型各组件的正确性。"""
print("\n" + "=" * 60)
print("运行组件测试...")
print("=" * 60)
batch, block_size = 2, 16
config = GPTConfig(
vocab_size=50, block_size=block_size,
n_layer=1, n_head=2, n_embd=32, dropout=0.0,
)
# 1. 测试单头注意力(AttentionHead)
causal_mask = torch.tril(torch.ones(1, 1, block_size, block_size))
head = AttentionHead(n_embd=32, head_size=16, causal_mask=causal_mask)
x = torch.randn(batch, block_size, 32)
out = head(x)
assert out.shape == (batch, block_size, 16), f"单头注意力输出形状错误: {out.shape}"
print("[OK] AttentionHead (单头注意力)")
# 2. 测试多头注意力(MultiHeadAttention)
mha = MultiHeadAttention(config)
out = mha(x)
assert out.shape == x.shape, f"多头注意力输出形状错误: {out.shape} != {x.shape}"
print("[OK] MultiHeadAttention (多头注意力)")
# 3. 测试多头拼接维度(核心验证!)
# 确认每个头输出 head_size 维,拼接后正好是 n_embd 维
head_outputs = [head(x) for head in mha.heads]
concat = torch.cat(head_outputs, dim=-1)
assert concat.shape[-1] == config.n_head * mha.head_size == config.n_embd, \
f"多头拼接维度错误: {concat.shape[-1]} != {config.n_embd}"
print(f"[OK] Multi-Head Concatenation ({config.n_head} 个头拼接 = {concat.shape[-1]} 维)")
# 4. 测试 MLP
mlp = MLP(config)
out = mlp(x)
assert out.shape == x.shape, f"MLP 输出形状错误: {out.shape} != {x.shape}"
print("[OK] MLP")
# 5. 测试 Block
block = Block(config)
out = block(x)
assert out.shape == x.shape, f"Block 输出形状错误: {out.shape} != {x.shape}"
print("[OK] Transformer Block")
# 6. 测试因果性(关键验证!)
# 检查第 t 个位置的输出是否不受未来位置影响:
# 改变输入最后一个 token,前 15 个位置的输出应该完全不变
mha.eval()
x1 = torch.randn(1, block_size, 32)
x2 = x1.clone()
x2[0, -1, :] = 999.0 # 大幅修改最后一个 token
with torch.no_grad():
o1 = mha(x1)
o2 = mha(x2)
assert torch.allclose(o1[0, :-1], o2[0, :-1], atol=1e-5), "因果性被破坏: 未来影响了过去"
print("[OK] Causal Property (未来不会影响过去)")
# 7. 测试完整 GPT 模型
gpt = GPT(config)
idx = torch.randint(0, 50, (batch, block_size))
logits, loss = gpt(idx, idx)
assert logits.shape == (batch, block_size, 50), f"模型输出形状错误: {logits.shape}"
assert loss is not None and loss.ndim == 0, "损失计算错误"
print("[OK] GPT Model Forward + Loss")
# 8. 测试推理模式(无 targets)
logits, loss = gpt(idx)
assert loss is None, "推理模式下 loss 应为 None"
print("[OK] GPT Inference Mode (loss=None)")
# 9. 测试梯度反向传播
loss = logits.sum()
loss.backward()
grad_ok = all(p.grad is not None for p in gpt.parameters() if p.requires_grad)
assert grad_ok, "梯度反向传播失败"
print("[OK] Backpropagation")
# 10. 测试生成
with torch.no_grad():
gen = gpt.generate(idx[:1], max_new_tokens=5, temperature=0.8)
assert gen.shape == (1, block_size + 5), f"生成序列形状错误: {gen.shape}"
print("[OK] Text Generation")
print("=" * 60)
print("所有组件测试通过!")
print("=" * 60)
# ============================================================================
# 入口
# ============================================================================
if __name__ == "__main__":
# 1. 运行组件测试
test_components()
# 2. 训练字符级 GPT 语言模型
train_char_gpt()
多头注意力分析
已经拿到完整中间结果(为了可复现,诊断脚本固定了随机种子,数值与直接运行略有不同,但流程完全一致)。下面以测试 2 为例(batch=1, seq=2, d_model=4, num_heads=2, d_k=2),从输入到输出逐步给出每个阶段的具体张量。
第 0 步:输入
Q、K:(1, 2, 4),两个 token,每 token 4 维:
token0: [1, 0, 0, 1]
token1: [0, 1, 1, 0]
V:(1, 2, 4)
token0: [1, 2, 3, 4]
token1: [5, 6, 7, 8]
mask:(1, 2, 2),全 1(不遮蔽)。
第 1 步:W_q / W_k / W_v 线性投影(形状不变 (1,2,4)→(1,2,4))
- 每个 Linear 的权重是
(4, 4),对每个 token 做y = xWᵀ + b,把 4 维向量重新混合。以 Q 为例:
W_q(Q):
token0: [-0.1740, 0.5039, -0.4816, -0.6769]
token1: [ 0.0544, 0.4242, -0.3578, -0.7550]
- V 投影后:
W_v(V):
token0: [-2.0068, 0.8546, 2.3432, -0.0966]
token1: [-4.4560, 0.8129, 4.6515, 1.1307]
第 2 步:.view(1, -1, 2, 2) 切头(形状 (1,2,4)→(1,2,2,2),语义 (batch, seq, head, d_k))
- 把每个 token 最后的 4 维从中间切开:前 2 维归头 0,后 2 维归头 1。以 Q 为例:
view 后 (seq, head, d_k):
token0: 头0=[-0.1740, 0.5039] 头1=[-0.4816, -0.6769]
token1: 头0=[ 0.0544, 0.4242] 头1=[-0.3578, -0.7550]
- V 同理:
token0: 头0=[-2.0068, 0.8546] 头1=[ 2.3432, -0.0966]
token1: 头0=[-4.4560, 0.8129] 头1=[ 4.6515, 1.1307]
第 3 步:.transpose(1, 2) 归拢各头(形状仍 (1,2,2,2),语义变为 (batch, head, seq, d_k))
- 只是重新排列维度,把"同一头的所有 token"放到一起,让每个头能独立做注意力。Q 变成:
头0:
token0: [-0.1740, 0.5039]
token1: [ 0.0544, 0.4242]
头1:
token0: [-0.4816, -0.6769]
token1: [-0.3578, -0.7550]
- K、V 同样重排(下面会用到的 V):
头0: token0=[-2.0068, 0.8546] token1=[-4.4560, 0.8129]
头1: token0=[ 2.3432, -0.0966] token1=[ 4.6515, 1.1307]
第 4 步:scores = Q·Kᵀ((1,2,2,2)×(1,2,2,2)→(1,2,2,2),语义 (batch, head, seq_q, seq_k))
- 对每个头独立做
(2×2)·(2×2)矩阵乘。以头 0 的score[0][0](token0 对 token0)为例:-0.1740×0.2639 + 0.5039×(-0.1146) = -0.1037 - 完整分数:
头0: [[-0.1037, 0.4678],
[-0.0343, 0.3979]]
头1: [[ 0.2084, 0.2756],
[ 0.2935, 0.2897]]
- 每行代表"某个 query 对两个 key 的相似度"。
第 5 步:/ math.sqrt(2) 缩放(形状不变)
- 防止点积过大导致 softmax 梯度消失。结果:
头0: [[-0.0733, 0.3308],
[-0.0242, 0.2814]]
头1: [[ 0.1474, 0.1949],
[ 0.2075, 0.2049]]
第 6 步:mask 处理(形状不变)
mask.unsqueeze(1)从(1,2,2)变成(1,1,2,2),广播到 4 维后与 scores 对齐。masked_fill(mask==0, -inf):本测试 mask 全 1,没有位置被遮蔽,数值不变。如果某位置是 0,会被填成-inf,softmax 后概率为 0。
第 7 步:F.softmax(..., dim=-1) 得到注意力权重((1,2,2,2))
- 对最后一维(每个 query 对所有 key)归一化,每行和为 1。以头 0 第一行为例:
softmax([-0.0733, 0.3308]) = [0.4003, 0.5997] - 完整权重:
头0: [[0.4003, 0.5997],
[0.4242, 0.5758]]
头1: [[0.4881, 0.5119],
[0.5007, 0.4993]]
- 这就是"每个 token 应该用多大比例去看其他 token"。
第 8 步:attention_weights · V((1,2,2,2)→(1,2,2,2),仍是 (batch, head, seq, d_k))
- 每个头用权重对 V 加权求和。以头 0 token0 为例:
0.4003×[-2.0068, 0.8546] + 0.5997×[-4.4560, 0.8129] = [-3.4755, 0.8296] - 完整结果:
头0: token0=[-3.4755, 0.8296] token1=[-3.4171, 0.8306]
头1: token0=[ 3.5248, 0.5316] token1=[ 3.4958, 0.5162]
第 9 步:.transpose(1, 2) 换回((batch, head, seq, d_k)→(batch, seq, head, d_k))
- 把"头"维度挪回第 2 维,每个 token 的两个头结果重新归到一起:
token0: 头0=[-3.4755, 0.8296] 头1=[3.5248, 0.5316]
token1: 头0=[-3.4171, 0.8306] 头1=[3.4958, 0.5162]
- 此时内存是非连续的,所以紧跟
.contiguous()复制成连续内存,为下一步view做准备。
第 10 步:.view(1, -1, 4) 合并头((1,2,2,2)→(1,2,4))
- 把每个 token 的
(头0 2维, 头1 2维)首尾拼接回 4 维,语义(batch, seq, d_model):
token0: [-3.4755, 0.8296, 3.5248, 0.5316]
token1: [-3.4171, 0.8306, 3.4958, 0.5162]
第 11 步:W_o 输出投影((1,2,4)→(1,2,4))
- 用另一个
Linear(4,4)混合两个头的信息,得到最终输出:
最终输出:
token0: [-1.4836, -1.3123, -2.0893, 1.1161]
token1: [-1.4640, -1.2986, -2.0568, 1.1086]
最终返回
output:(1, 2, 4),和输入形状一致,可直接喂给下一层。attention_weights:(1, 2, 2, 2),即(batch, heads, seq_q, seq_k),保留每个头各自的注意力分布,方便可视化或做分析。
形状变化总表
- 输入
(1,2,4)→ 投影(1,2,4)→ view(1,2,2,2)→ transpose(1,2,2,2)→ 注意力计算(1,2,2,2)(权重同形状)→ transpose+contiguous(1,2,2,2)→ view(1,2,4)→ W_o(1,2,4)。
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
def scaled_dot_product_attention(Q, K, V, mask=None):
d_k = Q.size(-1)
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)
if mask is not None:
scores = scores.masked_fill(mask == 0, float('-inf'))
attention_weights = F.softmax(scores, dim=-1)
output = torch.matmul(attention_weights, V)
return output, attention_weights
def test_scaled_dot_product_attention():
Q = torch.tensor([[[1.0, 0.0], [0.0, 1.0]]])
K = torch.tensor([[[1.0, 0.0], [0.0, 1.0]]])
V = torch.tensor([[[1.0, 2.0], [3.0, 4.0]]])
mask = torch.tensor([[[1, 1], [1, 1]]])
output, attention_weights = scaled_dot_product_attention(Q, K, V, mask)
print("Output:", output)
print("Attention Weights:", attention_weights)
print("Output shape:", output.shape)
print("Attention Weights shape:", attention_weights.shape)
class OneHeadAttention(nn.Module):
def __init__(self, d_model):
super(OneHeadAttention, self).__init__()
self.W_q = nn.Linear(d_model, d_model)
self.W_k = nn.Linear(d_model, d_model)
self.W_v = nn.Linear(d_model, d_model)
def forward(self, Q, K, V, mask=None):
Q = self.W_q(Q)
K = self.W_k(K)
V = self.W_v(V)
output, attention_weights = scaled_dot_product_attention(Q, K, V, mask)
return output, attention_weights
class MultiHeadAttention(nn.Module):
def __init__(self, d_model, num_heads):
super(MultiHeadAttention, self).__init__()
assert d_model % num_heads == 0
self.d_k = d_model // num_heads
self.num_heads = num_heads
self.W_q = nn.Linear(d_model, d_model)
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)
def forward(self, Q, K, V, mask=None):
batch_size = Q.size(0)
Q = self.W_q(Q).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2)
K = self.W_k(K).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2)
V = self.W_v(V).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2)
if mask is not None:
mask = mask.unsqueeze(1)
output, attention_weights = scaled_dot_product_attention(Q, K, V, mask)
output = output.transpose(1, 2).contiguous().view(batch_size, -1, self.num_heads * self.d_k)
output = self.W_o(output)
return output, attention_weights
def test_multi_head_attention():
d_model = 4
num_heads = 2
mha = MultiHeadAttention(d_model, num_heads)
Q = torch.tensor([[[1.0, 0.0, 0.0, 1.0], [0.0, 1.0, 1.0, 0.0]]])
K = torch.tensor([[[1.0, 0.0, 0.0, 1.0], [0.0, 1.0, 1.0, 0.0]]])
V = torch.tensor([[[1.0, 2.0, 3.0, 4.0], [5.0, 6.0, 7.0, 8.0]]])
mask = torch.tensor([[[1, 1], [1, 1]]])
output, attention_weights = mha(Q, K, V, mask)
print("MultiHead Output:", output)
print("MultiHead Attention Weights:", attention_weights)
print("MultiHead Output shape:", output.shape)
print("MultiHead Attention Weights shape:", attention_weights.shape)
if __name__ == "__main__":
test_scaled_dot_product_attention()
test_multi_head_attention()
# -*- coding: utf-8 -*-
"""debug_attention.py
逐步打印 MultiHeadAttention 前向传播中每个阶段的张量形状与数值,
方便理解 view / transpose / 注意力计算 / 合并头等每一步的变化。
用法:
python debug_attention.py # 运行内置示例(与 trans.py 测试 2 相同输入)
或在代码中调用 debug_attention(Q, K, V, mask, d_model, num_heads)
"""
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
def _show(name, tensor):
"""打印张量名称、形状和具体数值。"""
print(f"[{name}] shape={tuple(tensor.shape)}")
print(tensor)
def debug_attention(Q, K, V, mask=None, d_model=4, num_heads=2, seed=0):
"""按步骤打印多头注意力的中间结果。
参数:
Q, K, V : 形状为 (batch, seq, d_model) 的输入张量
mask : 可选,形状为 (batch, seq, seq),0 表示遮蔽
d_model : 模型维度
num_heads: 注意力头数
seed : 随机种子,固定后结果可复现
返回:
(output, attention_weights)
"""
assert d_model % num_heads == 0
d_k = d_model // num_heads
batch_size = Q.size(0)
# 固定随机种子,保证线性层初始化一致、结果可复现
torch.manual_seed(seed)
# 四个投影层,与 trans.py 中 MultiHeadAttention 的配置一致
W_q = nn.Linear(d_model, d_model)
W_k = nn.Linear(d_model, d_model)
W_v = nn.Linear(d_model, d_model)
W_o = nn.Linear(d_model, d_model)
print(f"=== 配置: d_model={d_model}, num_heads={num_heads}, d_k={d_k}, batch={batch_size}")
print(f"=== W_q.weight 形状: {tuple(W_q.weight.shape)}(按行切成 {num_heads} 块,每块就是每个头的 d_model->d_k 投影)") # ---- 第 0 步:输入 ----
_show("第0步 输入 Q", Q)
_show("第0步 输入 K", K)
_show("第0步 输入 V", V)
if mask is not None:
_show("第0步 输入 mask", mask)
# ---- 第 1 步:线性投影 (batch, seq, d_model) -> (batch, seq, d_model) ----
Qp = W_q(Q)
Kp = W_k(K)
Vp = W_v(V)
_show("第1步 W_q(Q) 线性投影后", Qp)
_show("第1步 W_k(K) 线性投影后", Kp)
_show("第1步 W_v(V) 线性投影后", Vp)
# ---- 第 2 步:view 切头 (batch, seq, d_model) -> (batch, seq, num_heads, d_k) ----
# 每个 token 的 d_model 维向量被均分成 num_heads 份,每份 d_k 维
Qv = Qp.view(batch_size, -1, num_heads, d_k)
Kv = Kp.view(batch_size, -1, num_heads, d_k)
Vv = Vp.view(batch_size, -1, num_heads, d_k)
_show("第2步 Q view 后 (batch, seq, head, d_k)", Qv)
_show("第2步 V view 后 (batch, seq, head, d_k)", Vv)
# ---- 第 3 步:transpose(1,2) (batch, seq, head, d_k) -> (batch, head, seq, d_k) ----
# 把"头"维度挪到序列长度之前,让每个头独立做注意力
Qt = Qv.transpose(1, 2)
Kt = Kv.transpose(1, 2)
Vt = Vv.transpose(1, 2)
_show("第3步 Q transpose 后 (batch, head, seq, d_k)", Qt)
_show("第3步 K transpose 后 (batch, head, seq, d_k)", Kt)
_show("第3步 V transpose 后 (batch, head, seq, d_k)", Vt)
# ---- 第 4 步:scores = Q*K^T (batch, head, seq_q, seq_k) ----
scores = torch.matmul(Qt, Kt.transpose(-2, -1))
_show("第4步 scores = Q*K^T(未缩放)", scores)
# ---- 第 5 步:除以 sqrt(d_k) 缩放,防止点积过大 ----
scores_scaled = scores / math.sqrt(d_k)
_show("第5步 scores / sqrt(d_k)", scores_scaled)
# ---- 第 6 步:mask 处理 ----
if mask is not None:
mask4 = mask.unsqueeze(1) # (batch, seq, seq) -> (batch, 1, seq, seq),广播到所有头
_show("第6步 mask unsqueeze(1) 后 (batch, 1, seq, seq)", mask4)
scores_scaled = scores_scaled.masked_fill(mask4 == 0, float("-inf"))
_show("第6步 mask 后 scores(-inf 处 softmax 后为 0)", scores_scaled) # ---- 第 7 步:softmax(dim=-1) 得到注意力权重,每行和为 1 ----
attention_weights = F.softmax(scores_scaled, dim=-1)
_show("第7步 attention_weights (batch, head, seq_q, seq_k)", attention_weights)
# ---- 第 8 步:权重对 V 加权求和 (batch, head, seq, d_k) ----
attn_out = torch.matmul(attention_weights, Vt)
_show("第8步 attention 输出 (batch, head, seq, d_k)", attn_out)
# ---- 第 9 步:transpose(1,2) 换回 (batch, seq, head, d_k),再 contiguous 保证内存连续 ----
merged = attn_out.transpose(1, 2).contiguous()
_show("第9步 transpose 回 (batch, seq, head, d_k) + contiguous", merged)
# ---- 第 10 步:view 合并头 (batch, seq, head*d_k) = (batch, seq, d_model) ----
merged2 = merged.view(batch_size, -1, num_heads * d_k)
_show("第10步 view 合并回 (batch, seq, d_model)", merged2)
# ---- 第 11 步:W_o 输出投影,混合各头信息 (batch, seq, d_model) ----
output = W_o(merged2)
_show("第11步 W_o 输出投影(最终输出)", output)
print(f"=== 最终: output 形状 {tuple(output.shape)}, attention_weights 形状 {tuple(attention_weights.shape)}")
return output, attention_weights
def main():
# 与 trans.py 中 test_multi_head_attention 相同的输入
Q = torch.tensor([[[1.0, 0.0, 0.0, 1.0], [0.0, 1.0, 1.0, 0.0]]])
K = torch.tensor([[[1.0, 0.0, 0.0, 1.0], [0.0, 1.0, 1.0, 0.0]]])
V = torch.tensor([[[1.0, 2.0, 3.0, 4.0], [5.0, 6.0, 7.0, 8.0]]])
mask = torch.tensor([[[1, 1], [1, 1]]])
debug_attention(Q, K, V, mask, d_model=4, num_heads=2, seed=0)
if __name__ == "__main__":
main()
更多推荐
所有评论(0)