从零实现 Attention:100 行代码的 Transformer 心脏

06-理论实践与脉络 核心 约 30 分钟 #实现#代码#PyTorch#attention 更新 2026-10-02
当前状态:未学
本文基于模型知识整理(生成时未联网核对),代码对照 Karpathy 的 nanoGPT / minGPT 与 PyTorch 官方教程复核。

一句话定义

把 kp-002/003/006 的公式变成 PyTorch 代码:缩放点积注意力→多头→Block→完整 GPT——100 行代码包含 Transformer 的全部心脏;亲手实现一遍是对所有概念(QKV 排列、因果掩码、残差流)的终极检验。

为什么重要

"能读懂公式"与"能写出正确实现"之间有一道实战鸿沟:张量维度排列、因果掩码位置、多头 reshape——每一处都是概念理解的试金石。Karpathy 的 nanoGPT 系列(视频+代码)证明:写完一遍后,读任何 LLM 论文的架构部分都不再吃力。

前置知识

kp-002/003/006(注意力/多头/Block)、Python/PyTorch 张量基础。

核心概念(完整实现框架)

import torch
import torch.nn as nn
import torch.nn.functional as F

class CausalSelfAttention(nn.Module):
    """多头因果自注意力 (kp-002 + kp-003)"""
    def __init__(self, d_model, n_head):
        super().__init__()
        self.n_head = n_head
        self.d_head = d_model // n_head
        # QKV 合并投影 (一次算三个, 高效)
        self.qkv = nn.Linear(d_model, 3 * d_model)
        self.proj = nn.Linear(d_model, d_model)  # W^O (kp-003)

    def forward(self, x):
        B, T, C = x.shape
        # (1) 投影并拆三份: (B,T,C) → 3×(B,T,heads,head_dim)
        q, k, v = self.qkv(x).split(C, dim=2)
        q = q.view(B, T, self.n_head, self.d_head).transpose(1, 2)
        k = k.view(B, T, self.n_head, self.d_head).transpose(1, 2)
        v = v.view(B, T, self.n_head, self.d_head).transpose(1, 2)
        # (2) 缩放点积 (kp-002): (B,h,T,T)
        att = (q @ k.transpose(-2, -1)) / (self.d_head ** 0.5)
        # (3) 因果掩码: 未来置-∞ (kp-002 误区)
        mask = torch.triu(torch.ones(T, T, device=x.device), diagonal=1)
        att = att.masked_fill(mask == 1, float('-inf'))
        att = F.softmax(att, dim=-1)
        # (4) 加权取V → 合并多头
        out = (att @ v).transpose(1, 2).contiguous().view(B, T, C)
        return self.proj(out)

class Block(nn.Module):
    """Pre-LN Block (kp-006 + kp-007)"""
    def __init__(self, d_model, n_head):
        super().__init__()
        self.ln1 = nn.LayerNorm(d_model)
        self.attn = CausalSelfAttention(d_model, n_head)
        self.ln2 = nn.LayerNorm(d_model)
        self.ffn = nn.Sequential(
            nn.Linear(d_model, 4 * d_model),
            nn.GELU(),                      # kp-008 激活
            nn.Linear(4 * d_model, d_model),
        )
    def forward(self, x):
        x = x + self.attn(self.ln1(x))      # 残差1 (Pre-LN)
        x = x + self.ffn(self.ln2(x))       # 残差2
        return x

原理与机制(代码中的概念检验点)

  • 检验点 1:QKV 的 view/transpose——split(C,dim=2) 把合并投影拆成 Q/K/V;.view(B,T,h,d).transpose(1,2) 把"每token的512维"变成"8个头各64维"——多头=换视角(kp-003),张量操作是概念的空间化。
  • 检验点 2:掩码的两种实现——triu(diagonal=1) 生成上三角(不含对角线,自己的 K 保留)+ masked_fill(-inf)——softmax 后未来为 0;忘了做 = 信息泄露(模型看到答案)。
  • 检验点 3:残差的简洁——x + attn(ln(x)) 一行就是 Pre-LN Block 的一半(kp-006/007);公式的代码化如此直接,侧面印证了"简单架构"的设计哲学。
  • 常见 bug 自检清单:

- * 代替 @(逐元素乘 vs 矩阵乘——kp-030 的 NumPy 教训); - 掩码对角线少一行(diagonal=0 vs 1); - transpose 后没 contiguous(view 报错); - 缩放因子用了 d_model 而非 d_head。

图示

数据流: (B,T,C)
 → qkv投影 → split+view+transpose → (B,h,T,d)
 → QKᵀ/√d → mask(-∞) → softmax → @V
 → transpose+view → (B,T,C) → W^O → +残差

检验: 每个张量操作 = 每个概念的空间化
mask = kp-002 的因果约束
view/transpose = kp-003 的多头切分
+残差 = kp-006 的信息主干道

直观类比

从零实现像"亲手拆装引擎":说明书(论文)读十遍,不如亲手拆一遍再装回去——每个螺丝(张量操作)的位置与作用都刻进手指记忆。nanoGPT 的 300 行 = 你读懂一切 LLM 代码的通行证。

实例或案例

  • Karpathy 的 nanoGPT:~300 行训练 GPT-2 级模型——从零实现的最小完备参考。
  • picoGPT:60 行 NumPy 推理(无 PyTorch)——推理侧的极简对照。
  • minGPT → nanoGPT → llm.c 的演进:同一逻辑从 Python 到 C++/CUDA 的下沉——实现语言的抽象层级。

常见误区

  • 误区一:"写完能跑=实现正确"。因果掩码漏一行/缩放因子错/多头 reshape 错——都能跑但效果差;必须用kp-029 的梯度检查或小规模已知解对照。
  • 误区二:"先写完整模型再调试"。从单头注意力开始逐组件验证(kp-029 的"小例验证"纪律)——逐步堆叠而非一次性全写。
  • 误区三:"自己实现用于生产"。生产用 HF transformers(优化 kernel/分布式/量化 kp-026/024);自己实现的目的是学习与原型——各归其位。

与其他知识点的关系

  • kp-002/003/006/008:全部概念的代码化。
  • kp-013/014:训练稳定性代码的实战面。
  • kp-033:nanoGPT/Karpathy 的教学贡献。

自测题

  1. 为什么 QKV 用一个 Linear(3*d) 而不是三个 Linear(d)?

答:数学等价(一次矩阵乘后 split)——一次 GEMM 比三次更高效(GPU 友好),是工程优化非概念变化。

  1. torch.triu(ones, diagonal=1) 的含义?

答:上三角(不含对角线)全 1——掩码矩阵:位置 t 只能看到 ≤t 的 Key(diagonal=1 表示排除自身之后的所有位)。

  1. 如何验证实现的因果注意力是正确的?

答:对同一个输入,逐位置检查"只看过去"(改变未来的 token 不影响过去的输出);或对比参考实现(nanoGPT/HF)的中间激活。

延伸阅读

  • Karpathy, "Let's build GPT from scratch"(YouTube,2 小时从零到训练)。
  • nanoGPT / minGPT GitHub(最可读的 Transformer 实现)。
  • The Annotated Transformer(Harvard NLP,论文逐段对照代码)。