使用注意力机制进行类GPT模型训练
一、相关原理
类GPT,把输入的token作为文章的开头,进行自回归输入,最终输出接下来的文本。相比Transformer(翻译),GPT只需要用到Decoder,相比之下比较好写[修改说明:原文保留,以下为扩充内容]。
1. 自回归生成的原理
什么是"自回归"?用一个具体的例子来说明。假设你正在写一段文字,目前已经写下了token序列:[“我”, “今天”, “去”, “超市”, “买”]——这几个字就像写作文写到一半,现在模型的任务是:接下来该写什么?
模型对着这5个token进行计算,最后输出一个概率分布,告诉我们:“了"这个词的概率最高。于是模型预测出了"了”。现在,完整的句子变成了[“我”, “今天”, “去”, “超市”, “买”, “了”]。然后呢?模型把这一整串——注意,是整个6个token——再喂给自己,继续预测第7个token。这次它可能预测出"一",那就变成7个token;再输入,预测出"支"…如此循环往复。
这就是"自回归"(Autoregressive)名字的由来:用自己上一步的输出作为下一步的输入,一步一步地把文本"生长"出来。每一步都依赖于前面已经生成的所有内容。
这也解释了为什么我们使用ChatGPT或者其他GPT类模型时,文本是一个字一个字地蹦出来的,而不是一次性生成整个段落。因为"一次性生成"在数学上就是不成立的——后面的词依赖前面的词,前面的词还没生成,后面的词无从谈起。每一个token的诞生,都必须等待它前面的所有token都已经就位。
这里有一个很重要的对比:翻译任务。在原始Transformer的翻译场景中,Encoder会先把整句源语言读完,得到一个完整的语义表示,然后Decoder在生成目标语言时,每一步既可以"看"已经生成的输出(用Masked Self-Attention),又可以"看"Encoder提供的源语言信息(用Cross-Attention)——这是一种"两边都看"的双向模式。而GPT不需要翻译,它只有一个方向:从前往后、单向生成。因此GPT只保留了Decoder,并且去掉了Cross-Attention——它没有源语言可以"参考",也不需要。这种单向生成的设计,让GPT天然适合文本续写、故事创作、代码补全等任务。
2. 因果掩码的直观解释
理解了自回归的原理之后,一个实际问题就出现了:训练的时候,我们明明有一整段完整的文本,所有token都是一次性摆在模型面前的。那怎么让模型"假装"它看不到后面的词,只能看到前面的词呢?
答案就是因果掩码(Causal Mask),也叫三角形掩码。
假设我们的输入序列是 [“我”, “今天”, “去”, “超市”, “买”, “了”, “一”, “支”, “笔”] 这9个token。在自注意力机制中,每个token都要和所有其他token计算注意力分数。用一个9×9的矩阵来表示这个"谁可以看到谁"的关系:
我 今天 去 超市 买 了 一 支 笔
我 ✓ ✗ ✗ ✗ ✗ ✗ ✗ ✗ ✗
今天 ✓ ✓ ✗ ✗ ✗ ✗ ✗ ✗ ✗
去 ✓ ✓ ✓ ✗ ✗ ✗ ✗ ✗ ✗
超市 ✓ ✓ ✓ ✓ ✗ ✗ ✗ ✗ ✗
买 ✓ ✓ ✓ ✓ ✓ ✗ ✗ ✗ ✗
了 ✓ ✓ ✓ ✓ ✓ ✓ ✗ ✗ ✗
一 ✓ ✓ ✓ ✓ ✓ ✓ ✓ ✗ ✗
支 ✓ ✓ ✓ ✓ ✓ ✓ ✓ ✓ ✗
笔 ✓ ✓ ✓ ✓ ✓ ✓ ✓ ✓ ✓
看第5行(对应token “买”):它只能"看"到前5个token——[“我”, “今天”, “去”, “超市”, “买”],后面的[“了”, “一”, “支”, “笔”]全部被遮住。这就形成了一个下三角(✓的区域)加一个上三角(✗的区域)的图案,所以叫"三角形掩码"。
在代码实现中,这个遮蔽是通过把上三角位置的注意力分数设为 -∞(实际代码中用 -1e4 近似)来实现的。经过Softmax之后,e^(-∞) = 0,那些位置的注意力权重就变成了0,达到了"看不到"的效果。
你可能会有疑问:这样做是不是在"限制"模型的能力?明明训练数据里有完整信息,为什么不给模型看?其实,这恰恰不是一种限制,而是一种对齐训练和推理的技巧。回想一下:在推理(实际生成文本)的时候,模型确实只能看到已经生成的前半部分,未来的词还没产生,当然看不到。如果训练时让模型能看到后面的词,那训练目标和推理环境就不一致了——模型在训练时学会了"抄答案",到了推理时却无答案可抄,表现必然大幅下降。因果掩码保证:训练时"模拟"了推理时的真实处境,两者保持一致。
这里还有一个重要的对比,值得反复理解:
-
Encoder中的自注意力:双向可见。比如BERT在做句子理解时,每个token可以看到它左边和右边的所有token——这对"理解"任务是必需的,因为要理解一个词往往需要上下文。比如"bank"这个词,不看后面的"river"还是"account",你就不知道它是"河岸"还是"银行"。
-
Decoder/GPT中的因果注意力:单向遮蔽。每个token只能看到它自己和它左边的token——这对"生成"任务是必需的,因为生成永远是从左到右、未知到已知的过程。生成时右边还没有内容,你不能参考不存在的未来。
两者并非谁优谁劣,而是服务于不同的目标:理解用双向,生成用单向。GPT选择了后者,因为它的使命不是理解一段话,而是延续一段话。
3. GPT与原始Transformer的关系
理解了自回归和因果掩码之后,我们来说清楚一个经常让人困惑的问题:GPT和原始Transformer到底有什么关系?
一句话总结:GPT的架构 = Transformer的Decoder部分,但去掉了Cross-Attention,只保留了Masked Self-Attention + FFN(前馈网络)。
原始Transformer(出自2017年的论文《Attention Is All You Need》)是一个Encoder-Decoder架构,设计用来做机器翻译:
- Encoder:读取源语言句子(如英文),通过双向自注意力理解整个句子的含义
- Decoder:一边用因果自注意力生成目标语言(如中文),一边用Cross-Attention参考Encoder的输出
GPT做的事情完全不同:它不做翻译,它是纯文本生成。输入是"文本的前半段",输出是"文本的后半段"。既然是同一个语言内部的续写,就不需要什么"源语言"的概念。所以:
- Cross-Attention直接不需要了——没有Encoder的输出可以"交叉注意"
- 只剩下Decoder中的Masked Self-Attention + FFN,一层一层堆叠起来
这就是为什么说GPT"只用到了Decoder"。
那OpenAI为什么要选择这条路线呢?这背后有很深的洞察:
解码器架构天然适合生成任务。生成永远是单向的——从已有内容推导后续内容。Masked Self-Attention正好完美地建模了这个过程。更重要的是,“预测下一个词"这个训练目标极其简单却极其强大——它不需要任何人工标注。你随便找一段文本,把它切成"前N个词"和"第N+1个词”,就构成了一条训练样本。这意味着你可以在互联网级别的海量无标注文本上进行预训练——这正是GPT系列成功的核心原因。
从GPT-1(2018年,1.17亿参数)到GPT-4(2023年,传闻1.76万亿参数),OpenAI走的就是一条Scaling路线:
- 架构上变化其实不大——核心还是Transformer Decoder的堆叠
- 变化的是规模:模型越来越大(更多层、更宽维度、更多注意力头)、数据越来越多(从几千本书到整个互联网)、算力越来越猛(从几块GPU到几万块GPU)
正如OpenAI在2020年那篇著名的Scaling Laws论文中所证明的:语言模型的性能与模型大小、数据量、算力呈幂律关系——只要持续加大投入,loss就会持续降低。本文中的小模型训练过程,其实就是这个Scaling Law在微观尺度上的一个缩影。
二、具体效果
现在一共训练了55万个Batch,取其中一些输出作为训练效果的体现,具体如下:
·batch-14000; loss-5.2

词汇表还没有多少,正在学习基本语句
·batch-25000; loss-3.5
句子结构基本正确,但是逻辑欠缺,出现一些不明所以的句子;对部分词语理解有误
·batch-48000; loss-3.0

对部分词的理解欠缺,句子上下文衔接不连贯
·batch-69000; loss-2.5

逻辑转换莫名其妙,出现不明所以的人物,词性未完全理解
·batch-97000; loss-2.2

基本语句流畅,句间逻辑有很大问题,并且重复的词比较多
·batch-117000; loss-2.0

句子语法正确,词义理解基本正确,重复的词汇减少,但是句间逻辑依然有比较大的问题
·batch-197000; loss-1.75

还在学习句子之间的逻辑,明显比上面好,但是出现莫名其妙的转折点
·batch-384000; loss-1.55

句子之间的逻辑明显好很多,形容词增加,句子成分更加复杂
·batch-425000; loss-1.5

生成的故事已经比较有逻辑了,中间可能有些断断续续,突然出现了一些莫名其妙的内容,但是至少已经很连贯了
·batch-514000; loss-1.45

对于长故事的创造力很强,但是句子之间还是缺乏逻辑,以及连接的时候会有很多问题。GPT还在天马行空的想象吧!
对上面的内容进行一个总结:
1、loss的折线图如下所示,逐渐趋缓并且达到一个瓶颈期

2、学习内容、产出文本解读
在多次实验中,生成参数保持一致,唯一的变量是所使用的模型。根据 Transformer 的结构特点,每个注意力头(Multi-Head Attention)都会捕捉不同层面的上下文关系。从生成结果可以看出,GPT 已经能够较好地理解词语的基本含义以及上下文之间的常见组合。然而,它在深入理解整体语义、以及在自回归生成过程中有效记忆较远的上下文内容方面,仍存在一定的不足。
深入解读:从训练过程中我们看到了什么
把上面的所有训练快照串起来看,我们能发现很多有意思的规律。
一、Loss曲线的幂律关系
从batch-14000(loss 5.2)到batch-514000(loss 1.45),loss一直在下降,但下降的速度越来越慢。这在坐标系中形成了一条典型的"先陡后平"的曲线——这正是**Scaling Law(规模定律)**的微观体现。
Scaling Law告诉我们:loss与训练数据量之间存在幂律关系,即 loss ∝ (数据量)^(-α),其中α是一个正指数。翻译成直观的理解就是:数据量翻倍,loss并不是等比例地降,而是以一个递减的速率在降。从5.2降到2.0相对容易(只需要大约12万batch),但从2.0降到1.45却需要大约40万batch——每往下压零点几个点,都要付出比之前大得多的数据代价。
这也意味着,本文的训练已经明显进入了瓶颈期。如果要继续降低loss,靠现在的模型规模已经不够了——可能需要从以下方向突破:
- 增大模型容量:增加d_model(嵌入维度)、n_heads(注意力头数)或n_layers(层数),让模型有更多的参数来吸收更复杂的模式
- 增加训练数据:更多的文本、更丰富的领域覆盖,让模型见过更多样的语言现象
- 扩展上下文长度:当前max_seq_len为128,很多长距离依赖关系可能因此被截断
二、模型"学会语言"的阶段性路径
更有趣的是,如果仔细阅读每个batch阶段的生成输出和评价,我们可以清晰地看到模型学习语言的阶段性路径——这不是这个模型独有的,而是几乎所有语言模型训练过程中都会经历的通用规律:
-
第一阶段:学词汇(batch ~14k,loss 5.2)——词汇表还没有多少,模型刚从随机初始化开始,正在学习最基础的东西:哪些token经常出现在一起,基本的词与词的搭配。这个阶段输出基本是乱码或极短的片段
-
第二阶段:学语法(batch ~25k-48k,loss 3.5-3.0)——句子结构开始正确,说明模型已经掌握了基本的语法规则(主谓宾结构、虚词用法等)。但语义层面还很脆弱,会出现一些"语法正确但语义不通"的句子,对部分词语的理解也时常跑偏
-
第三阶段:学局部逻辑(batch ~69k-117k,loss 2.5-2.0)——句子级别已经比较流畅,词汇理解基本正确,重复现象减少。但句子之间的逻辑衔接依然成问题——会突然引入莫名其妙的人物或出现不合逻辑的转折。这说明模型还在学习"前后两句话应该有什么样的关系"
-
第四阶段:学篇章连贯(batch ~197k+,loss 1.75及以下)——长故事有了一定的创造力,形容词增多,句子成分变复杂。模型开始理解更宏观的叙事结构和篇章逻辑,但仍然会出现"天马行空"式的跳跃——想象力和逻辑性之间的平衡还需要更多训练来打磨
这个学习路径揭示了一个深刻的道理:语言模型的训练,本质上就是从"字面"到"语义"、从"局部"到"全局"的渐进式抽象过程。先学会词是什么,再学会怎么组成句子,然后学会句子之间怎么连接,最后学会怎么构建一个完整的、连贯的篇章。每一步都建立在前一步的基础之上,不能跳跃。
三、实现代码(仅展示部分重要模块)
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
class MultiHeadAttention(nn.Module):
def __init__(self, d_model, n_heads, dropout=0.1):
super().__init__()
assert d_model % n_heads == 0
self.d_model = d_model
self.n_heads = n_heads
self.d_k = d_model // n_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)
self.dropout = nn.Dropout(dropout)
def forward(self, x, mask=None):
B, T, C = x.shape
Q = (
self.W_q(x)
.view(B, T, self.n_heads, self.d_k)
.transpose(1, 2)
)
K = (
self.W_k(x)
.view(B, T, self.n_heads, self.d_k)
.transpose(1, 2)
)
V = (
self.W_v(x)
.view(B, T, self.n_heads, self.d_k)
.transpose(1, 2)
)
# Attention Score
scores = (Q @ K.transpose(-2, -1)) / math.sqrt(self.d_k)
# Causal Mask
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e4)
attn = F.softmax(scores.float(), dim=-1).type_as(scores)
attn = self.dropout(attn)
output = (
(attn @ V)
.transpose(1, 2)
.contiguous()
.view(B, T, C)
)
return self.W_o(output)
class Model(nn.Module):
def __init__(
self,
vocab_size,
d_model=512,
n_heads=8,
num_layers=12,
max_seq_len=128,
d_ff=2048,
dropout=0.1
):
super().__init__()
self.d_model = d_model
self.max_seq_len = max_seq_len
# Token Embedding
self.embedding = nn.Embedding(
vocab_size,
d_model
)
# Position Embedding
self.pos_embedding = nn.Embedding(
max_seq_len,
d_model
)
self.dropout = nn.Dropout(dropout)
# Transformer Blocks
self.blocks = nn.ModuleList([
Block(
d_model,
n_heads,
d_ff,
dropout
)
for _ in range(num_layers)
])
self.ln_f = nn.LayerNorm(d_model)
# Output Projection
self.head = nn.Linear(
d_model,
vocab_size,
bias=False
)
# Weight Tying
self.embedding.weight = self.head.weight
self.apply(self._init_weights)
def _init_weights(self, module):
if isinstance(module, nn.Linear):
torch.nn.init.normal_(
module.weight,
mean=0.0,
std=0.02
)
if module.bias is not None:
torch.nn.init.zeros_(module.bias)
elif isinstance(module, nn.Embedding):
torch.nn.init.normal_(
module.weight,
mean=0.0,
std=0.02
)
elif isinstance(module, nn.LayerNorm):
torch.nn.init.zeros_(module.bias)
torch.nn.init.ones_(
module.weight
)
def forward(self, idx, targets=None):
B, T = idx.shape
device = idx.device
# Token embedding
tok_emb = self.embedding(idx)
# Position embedding
pos = torch.arange(
0,
T,
dtype=torch.long,
device=device
).unsqueeze(0)
pos_emb = self.pos_embedding(pos)
# Combine embeddings
x = self.dropout(
tok_emb + pos_emb
)
# Causal Attention Mask
mask = torch.tril(
torch.ones(
T,
T,
device=device
)
).view(
1,
1,
T,
T
)
# Transformer layers
for block in self.blocks:
x = block(
x,
mask
)
x = self.ln_f(x)
logits = self.head(x)
loss = None
if targets is not None:
loss = F.cross_entropy(
logits.view(-1, logits.size(-1)),
targets.view(-1)
)
# 防止异常 loss
loss = torch.clamp(
loss,
0,
15
)
return logits, loss