Skip to content

Repository files navigation

TorchAttention:从零实现基础 Transformer

这个项目通过手写 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。

Transformer 配置

配置使用 @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__ 进行检查。

Feed-Forward Network

每个 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

从逐头循环到融合 QKV

最直观的实现可以为每个头创建一个 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 并行计算。

QKV 张量变形

假设:

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 中并行计算。

Scaled Dot-Product Attention

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,
)

构造过程:

  1. torch.ones(..., dtype=torch.bool) 创建全 True 矩阵。
  2. torch.tril() 保留主对角线和下方元素,将上三角变为 False
  3. view(1, 1, T, T) 增加 batch 和 head 广播维度。
  4. register_buffer() 将其注册为非训练状态,使它能随模型移动到 CPU 或 GPU。
  5. persistent=False 表示不将可重新生成的 mask 保存进 state_dict

运行时只截取当前序列需要的部分:

causal_mask = self.causal_mask[:, :, :seq_length, :seq_length]

mask 的形状是:

(1, 1, T, T)

它会自动广播到:

(B, H, T, T)

masked_fill 与 finfo

在 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,因此通常不需要构造完整的三角形矩阵,但因果约束仍然存在。

TransformerBlock

当前 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

为什么调用 module(x)

使用 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,序列长度增加时显存占用按 增长。
  • 当前没有 KV Cache,推理时会重复计算历史 token 的 K/V。
  • test_attn.pytest_model.py 尚待补充正式测试。

下一步

建议按以下顺序继续:

  1. 完成 RoPE,并只将旋转位置编码应用到 Q 和 K。
  2. 将手写 attention 与 torch.nn.functional.scaled_dot_product_attention 做数值对比。
  3. 添加 padding mask 和正式单元测试。
  4. 实现 KV Cache,区分 prefill 和单 token decode。
  5. 实现分块 attention 和 online softmax,逐步过渡到 FlashAttention。

About

No description, website, or topics provided.

Resources

Stars

Watchers

Forks

Releases

Packages

Contributors

Languages