GPT:从因果注意力到下一个 token 预测

用一段最小 PyTorch 实现理解 GPT 的输入、掩码、训练目标与生成过程。

先看它在做什么

给定 token 序列 x₀, x₁, …, xₜ,GPT 预测下一个 token:p(xₜ₊₁ | x₀, …, xₜ)。例如输入“今天下雨,我带了”,模型会为“伞”等候选词分配概率。生成时选出一个 token,接到序列末尾,再预测下一个;反复执行就得到一段文字。token 是分词器定义的单位,并不总是一个汉字或一个完整单词。

从原始文字到模型输出有一条清楚的路径:分词器把文字变成整数 ID,嵌入表把 ID 变成向量;若干 Transformer 块混合历史信息;最后线性层为词表中每个 token 给出一个未归一化分数(logit)。对这些分数做 softmax 才是概率。模型并不直接“写出”字符,选中的 token 还需要由分词器解码回文字。

这里的 GPT 指一类 decoder-only Transformer 自回归语言模型。2018 年的原始 GPT(后称 GPT-1)使用 Transformer 解码器式堆叠,先做语言模型预训练,再针对具体任务做有监督微调。后续 GPT 系列在规模、数据和使用方式上继续演变;下面讲的是家族共有的核心计算,并非声称所有版本都沿用 GPT-1 的训练流程。原始论文见 Radford 等,Improving Language Understanding by Generative Pre-Training。

为什么只能看左边

普通自注意力可让位置 t 读取整段输入。语言模型训练时,如果它看见位置 t+1 的正确 token,再预测这个 token,答案就泄漏了。因此在注意力分数矩阵上加因果掩码:第 t 行只保留列 0…t,未来列设为负无穷,softmax 后其权重为零。矩阵是下三角形;这正是“自回归”的结构约束。

每层把隐藏状态线性投影成查询 Q、键 K、值 V。一个注意力头计算 softmax(QKᵀ / √d_head + mask)V;多头注意力把不同头的输出拼接后再投影。随后是逐位置前馈网络,实际模型还会使用残差连接和归一化。token 嵌入本身不提供顺序信息,所以还需要位置表示。下面选用可学习绝对位置嵌入来说明原理;不同 GPT 变体的位置方案可能不同。

这里的“decoder-only”说的是网络结构:它没有单独的编码器输出供交叉注意力读取。它与机器翻译中“编码器读源语言、解码器写目标语言”的完整 encoder–decoder 架构不同。一个 GPT 块内部仍有自注意力和前馈网络;“decoder”一词并不表示它只能在生成时工作,训练时它会并行处理一整段序列。

一段可读的最小实现

下面的输入 tokens 形状是 [B, T],输出 logits 为 [B, T, V]:B 是批量大小,T 是长度,V 是词表大小。代码保留一个注意力块、前馈层、残差与归一化,便于看清数据流;它不是完整的 GPT-1 复现。

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

class TinyGPT(nn.Module):
    def __init__(self, vocab_size, width=128, heads=4, max_len=512):
        super().__init__()
        assert width % heads == 0
        self.heads = heads
        self.token = nn.Embedding(vocab_size, width)
        self.pos = nn.Embedding(max_len, width)
        self.norm1 = nn.LayerNorm(width)
        self.qkv = nn.Linear(width, 3 * width)
        self.proj = nn.Linear(width, width)
        self.norm2 = nn.LayerNorm(width)
        self.ffn = nn.Sequential(nn.Linear(width, 4 * width), nn.GELU(),
                                 nn.Linear(4 * width, width))
        self.final_norm = nn.LayerNorm(width)
        self.lm_head = nn.Linear(width, vocab_size)

    def forward(self, tokens):
        batch, length = tokens.shape
        assert length <= self.pos.num_embeddings
        x = self.token(tokens) + self.pos(torch.arange(length, device=tokens.device))
        h = self.norm1(x)
        q, k, v = self.qkv(h).chunk(3, dim=-1)
        def split(t):
            return t.reshape(batch, length, self.heads, -1).transpose(1, 2)
        q, k, v = map(split, (q, k, v))  # [B, H, T, d_head]
        scores = q @ k.transpose(-2, -1) / (q.size(-1) ** 0.5)
        future = torch.ones(length, length, device=tokens.device, dtype=torch.bool).triu(1)
        weights = scores.masked_fill(future, float('-inf')).softmax(dim=-1)
        attended = (weights @ v).transpose(1, 2).reshape(batch, length, -1)
        x = x + self.proj(attended)
        x = x + self.ffn(self.norm2(x))
        return self.lm_head(self.final_norm(x))

# 同一段文本的相邻位置形成输入与标签;tokens 至少有 2 个位置。
inputs, targets = tokens[:, :-1], tokens[:, 1:]
logits = model(inputs)                         # [B, T-1, V]
loss = F.cross_entropy(logits.reshape(-1, logits.size(-1)), targets.reshape(-1))

最后三行假定已准备好整数 token 张量 tokens 和 model = TinyGPT(vocab_size)。例子没有加入 padding;真实批次若有填充位置,损失与注意力还需正确处理 padding。代码中的致密注意力矩阵大小为 T × T,长序列会带来显存与计算开销。

留意 targets = tokens[:, 1:] 的错位:输入的第一个位置看见 x₀,其标签是 x₁;输入的最后一个位置看见到 xₜ₋₁,其标签是 xₜ。不要把同一位置的输入和标签直接对齐,否则任务会变成复制已见 token。示例仅有一个块,也省略了 dropout、训练循环与更高效的注意力实现,目的是突出这个预测关系。

Teacher forcing 与生成

训练时把整段正确文本一次送入模型,用第 t 个位置的输出预测第 t+1 个 token。这叫 teacher forcing:上下文来自真实文本,而非模型前一步的猜测。因果掩码保证虽然所有位置并行计算,每个位置仍只读到自己的历史。训练目标是所有目标位置的交叉熵均值,也就是最大化正确下一个 token 的对数概率。

推理时没有未来的正确文本,只能将刚生成的 token 追加到上下文,再做下一步预测。采样温度、top-k 等策略改变如何从预测分布选 token,不改变模型的因果约束。若只记一个要点:GPT 的“生成”来自重复做下一个 token 预测;掩码使训练时的并行预测与推理时的逐步生成遵守同一信息边界。

训练和推理的另一个差别是上下文来源:训练始终基于真实前缀,推理时则可能基于模型自己先前生成的不准确 token。因此训练损失低不保证长段生成始终可靠。模型也只会利用其上下文窗口内的信息;超出窗口的早期 token,需要截断、压缩或借助外部记忆等额外方法处理。

原论文