这个项目通过手写 PyTorch 代码,逐步理解 decoder-only Transformer 的核心结构,并为后续实现 KV Cache 和 FlashAttention 做准备。
当前已经完成一个基础版本,包含:
- Token Embedding
- Learned Position Embedding
- 融合 QKV 的 Multi-Head Self-Attention
- Causal Mask
- Attention Dropout
- Feed-Forward Network
- Residual Connection
- LayerNorm
- Language Model Head
RoPE 尚未完成,当前默认使用 learned position embedding。
| 阶段 | 内容 | 目标 |
|---|---|---|
| 1 | 基础 Transformer | 理解 QKV、多头注意力、因果掩码和残差结构 |
| 2 | RoPE | 将位置信息应用到 Q 和 K |
| 3 | KV Cache | 推理时复用历史 K/V,避免重复计算 |
| 4 | FlashAttention | 分块计算 attention,减少显存读写和中间矩阵 |
| 文件 | 作用 |
|---|---|
config.py |
Transformer 超参数配置 |
transformer.py |
FFN、Multi-Head Attention、TransformerBlock 和完整模型 |
rope.py |
RoPE,实现中 |
attention.py |
独立 attention 实验代码 |
test_attn.py |
Attention 测试文件 |
test_model.py |
完整模型测试文件 |
项目使用 Python 3.12+ 和 PyTorch,通过 uv 管理依赖:
uv sync运行模型:
uv run python transformer.py默认配置接近 GPT-2 Small 的规模,在 CPU 或显存较小的 GPU 上运行可能比较慢。调试时建议使用小配置,例如:
config = TransformerConfig(
vocab_size=100,
d_model=64,
n_heads=4,
d_ff=256,
n_layer=2,
max_seq_len=32,
pos_emb="learned",
)输入是一组 token ID:
(batch_size, seq_length)
完整的数据流为:
token IDs
│
├── Token Embedding
└── Position Embedding
│
▼
(batch, seq, d_model)
│
▼
TransformerBlock × N
│
▼
Final LayerNorm
│
▼
LM Head Linear
│
▼
(batch, seq, vocab_size)
最终输出的最后一个维度是 vocab_size,因为每个位置都需要为词表中的每个 token 产生一个 logit。
配置使用 @dataclass 管理:
@dataclass
class TransformerConfig:
vocab_size: int = 50257
d_model: int = 768
n_heads: int = 12
d_ff: int = 4 * d_model
n_layer: int = 12
max_seq_len: int = 1024每个注意力头的维度由以下属性计算:
@property
def d_head(self) -> int:
return self.d_model // self.n_heads例如:
d_model = 768
n_heads = 12
d_head = 768 / 12 = 64
d_model 必须能够被 n_heads 整除,因此配置中使用 __post_init__ 进行检查。
每个 TransformerBlock 中包含一个逐 token 的 FFN:
d_model → d_ff → GELU → d_model
代码结构为:
self.fc1 = nn.Linear(config.d_model, config.d_ff)
self.fc2 = nn.Linear(config.d_ff, config.d_model)
self.activation = nn.GELU()FFN 不会混合不同 token,它只对每个 token 的隐藏向量独立执行相同的非线性变换。不同 token 之间的信息交换发生在 attention 中。
标准多头注意力中,每个头都会看到完整的 d_model 维词向量,然后将它投影到自己的 d_head 维子空间。
对于第 i 个头:
X: (batch, seq, d_model)
Wq_i: (d_model, d_head)
Q_i: (batch, seq, d_head)
因此,不应该先把原始词向量固定切成多段,让每个头只能看到其中一段。正确顺序是:
完整词向量 → 线性投影 → 拆分多个 head
最直观的实现可以为每个头创建一个 SelfAttn,然后遍历:
attn_outputs = [head(x) for head in self.attn_heads]这里的 head 不是整数,而是 ModuleList 中保存的 SelfAttn 模块对象,因此 head(x) 会通过 nn.Module.__call__() 调用该模块的 forward()。
这种实现便于理解,但每个头都会产生独立的 Python 调用和矩阵乘法。当前实现将所有头的参数合并为一个大矩阵:
self.W_qkv = nn.Linear(
config.d_model,
3 * config.d_model,
bias=False,
)一次投影同时得到所有头的 Q、K、V:
x: (batch, seq, d_model)
qkv: (batch, seq, 3 × d_model)
这种表示在数学上等价于多个独立投影,但更适合 GPU 并行计算。
假设:
batch_size = B
seq_length = T
d_model = C
n_heads = H
d_head = D
C = H × D
融合投影后:
qkv = self.W_qkv(x)形状为:
(B, T, 3C)
然后 reshape:
qkv = qkv.view(B, T, 3, H, D)形状变为:
(batch, seq, QKV, heads, head_dim)
接下来调整维度顺序:
qkv = qkv.permute(2, 0, 3, 1, 4)permute(2, 0, 3, 1, 4) 表示按照原张量的第 2、0、3、1、4 维重新排列:
调整前:(B, T, 3, H, D)
调整后:(3, B, H, T, D)
将 QKV 维度移到最前面后,可以沿第 0 维拆开:
q, k, v = qkv.unbind(dim=0)得到:
q: (B, H, T, D)
k: (B, H, T, D)
v: (B, H, T, D)
这使 batch 和 head 都成为批量维度,所有 head 可以在一次 torch.matmul 中并行计算。
Attention 的计算公式是:
Attention(Q, K, V) = softmax(QKᵀ / sqrt(d_head))V
首先计算 Q 和 K 的相似度:
attn_scores = torch.matmul(
q,
k.transpose(-2, -1),
) / (self.d_head ** 0.5)形状变化:
q: (B, H, T, D)
k.transpose(-2, -1): (B, H, D, T)
attn_scores: (B, H, T, T)
除以 sqrt(d_head) 是为了避免维度较大时点积数值过大,使 softmax 进入梯度很小的饱和区域。
经过 softmax 和 V 加权后:
attn_weights = torch.softmax(attn_scores, dim=-1)
attn_output = torch.matmul(attn_weights, v)得到:
attn_output: (B, H, T, D)
最后将所有 head 合并:
attn_output = attn_output.transpose(1, 2).contiguous()
attn_output = attn_output.view(B, T, C)
output = self.W_o(attn_output)decoder-only 语言模型一次处理完整训练序列。如果不使用因果掩码,较早位置可以直接看到后面的正确 token,从而发生答案泄漏。
因果注意力要求第 i 个 token 只能关注位置 0...i:
True False False False
True True False False
True True True False
True True True True
其中:
True → 可以关注
False → 未来位置,不允许关注
当前实现预先创建最大序列长度的 mask:
self.register_buffer(
"causal_mask",
torch.tril(
torch.ones(
config.max_seq_len,
config.max_seq_len,
dtype=torch.bool,
)
).view(1, 1, config.max_seq_len, config.max_seq_len),
persistent=False,
)构造过程:
torch.ones(..., dtype=torch.bool)创建全True矩阵。torch.tril()保留主对角线和下方元素,将上三角变为False。view(1, 1, T, T)增加 batch 和 head 广播维度。register_buffer()将其注册为非训练状态,使它能随模型移动到 CPU 或 GPU。persistent=False表示不将可重新生成的 mask 保存进state_dict。
运行时只截取当前序列需要的部分:
causal_mask = self.causal_mask[:, :, :seq_length, :seq_length]mask 的形状是:
(1, 1, T, T)
它会自动广播到:
(B, H, T, T)
在 softmax 之前屏蔽未来位置:
attn_scores = attn_scores.masked_fill(
~causal_mask,
torch.finfo(attn_scores.dtype).min,
)masked_fill(mask, value) 会将 mask=True 的位置替换为指定值。原 causal mask 中 True 表示允许关注,所以先通过 ~causal_mask 进行布尔取反,得到所有需要屏蔽的位置。
torch.finfo(dtype) 用于查询某种浮点类型的信息,其中 .min 是该类型能够表示的最小有限值:
float32 → 约 -3.4 × 10^38
float16 → -65504
bfloat16 → 约 -3.4 × 10^38
未来位置被替换为极大负数后:
softmax(极大负数) ≈ 0
因此这些位置不会参与后面的 V 加权求和。不能将未来位置简单填成 0,因为 softmax(0) 仍然会产生非零概率。
训练时使用因果掩码主要是防止答案泄漏。推理时未来 token 尚不存在,但标准 GPT 仍保持因果约束,主要原因是:
- 保持训练和推理的计算方式一致。
- 保证每个位置的隐藏状态只依赖当前及之前的 token。
- 保证历史 token 的 K/V 不会随着新 token 到来而改变。
- 允许 KV Cache 安全复用历史 K/V。
在没有 KV Cache、每次重新计算完整前缀的实现中,仍应使用相同的 causal mask。使用 KV Cache 单 token 解码时,缓存中本身不存在未来 token,因此通常不需要构造完整的三角形矩阵,但因果约束仍然存在。
当前 Block 使用 Post-LayerNorm:
attn_output = self.attn(x)
x = self.LN1(x + attn_output)
ffn_output = self.ffn(x)
x = self.LN2(x + ffn_output)结构为:
x ── Attention ── Add ── LayerNorm
│
└── FFN ── Add ── LayerNorm
LayerNorm 的归一化维度是 d_model。如果输入形状为:
(batch, seq, d_model)
那么 x.size(-1) 就是最后一个维度 d_model。LayerNorm 会独立归一化每个 token 的隐藏向量。
LayerNorm、Attention 和 FFN 都必须在 __init__() 中创建并注册,不能在 forward() 中临时创建,否则参数无法被优化器持续训练,也无法正确进入模型的 state_dict。
使用 PyTorch 模块时推荐:
attn_output = self.attn(x)而不是直接调用:
attn_output = self.attn.forward(x)因为 nn.Module 实现了 __call__():
self.attn(x)
↓
nn.Module.__call__()
↓
执行 hooks 等模块机制
↓
self.attn.forward(x)
定义模块时实现 forward(),使用模块时调用 module(...)。
- RoPE 尚未完成,目前应使用
pos_emb="learned"。 - 当前没有实现 padding mask,带 padding 的 batch 还需要额外屏蔽 padding token。
resid_pdrop已在配置中声明,但尚未应用到残差分支。- 当前会显式创建
(B, H, T, T)attention score,序列长度增加时显存占用按T²增长。 - 当前没有 KV Cache,推理时会重复计算历史 token 的 K/V。
test_attn.py和test_model.py尚待补充正式测试。
建议按以下顺序继续:
- 完成 RoPE,并只将旋转位置编码应用到 Q 和 K。
- 将手写 attention 与
torch.nn.functional.scaled_dot_product_attention做数值对比。 - 添加 padding mask 和正式单元测试。
- 实现 KV Cache,区分 prefill 和单 token decode。
- 实现分块 attention 和 online softmax,逐步过渡到 FlashAttention。