从零复现GPT-2 124M:手写Transformer语言模型的完整实战记录
最近花了两周时间,从零开始复现了一个GPT-2 124M参数的语言模型,包括数据准备、模型搭建、训练、生成和评估全流程。这个量级的模型卡在消费级硬件和学术研究之间的微妙平衡点:它足够大,能让你真正理解大模型训练中的显存管理、分布式通信、学习率调度等工程细节,又足够小,不至于让一张A100跑半个月才看到结果。这篇文章把我的完整过程和踩坑记录整理出来,给准备动手实践大模型训练的朋友一个可参考的路线图。
先说清楚复现目标。GPT-2 124M是OpenAI GPT-2系列里最小的版本,12层Transformer decoder,12个注意力头,768维隐藏层,词汇表大小50257,最大序列长度1024。这个结构的每一层都被后来的GPT-3、LLaMA等模型继承或改造,所以搞清楚它的实现,几乎等于拿到了现代大语言模型的通用钥匙。适合的读者是:会Python和基础PyTorch,但没完整写过Transformer训练流程的人;或者已经用HuggingFace跑过推理,但想看看底层到底怎么工作的人。
我的复现路线是纯手工实现模型结构(不依赖transformers库的模型代码),用tiktoken做分词,用公开文本数据集做训练。整个流程分为数据、模型、训练、生成四块,下面按实际操作顺序展开。
1. 复现前的准备:硬件评估与技术选型
1.1 硬件门槛与显存估算
很多人第一反应是“124M参数,好像也不大”,但真正跑起来才发现,模型参数只是显存消耗的一部分。以我的实际配置为例:一张24GB显存的RTX 3090,batch size设8,序列长度1024,混合精度训练,显存占用约20GB,勉强能跑。
显存开销的大头有四块:模型参数(124M × 4字节 ≈ 500MB fp32,混合精度下权重和优化器状态会翻到约2GB)、梯度(和参数等量)、优化器状态(AdamW需要一阶动量、二阶动量,再加权重本身,约是参数量的8倍)、激活值(前向传播过程中需要保留的中间结果,这个和batch size、序列长度成正比)。算下来一个124M模型的理论显存底线是:
- 纯FP32推理:约0.5GB
- 混合精度训练:约6GB起步
- 加上激活值和通信缓冲:实际训练建议至少16GB
如果显存不够,两个变通方案我实测过:一是降低batch size配合梯度累积,二是在序列长度上做截断(比如从1024降到512,显存几乎减半),但后者会影响模型对长距离依赖的建模能力,属于不得已的选择。
1.2 技术栈与关键依赖版本
我的环境是Ubuntu 20.04 + Python 3.10 + CUDA 11.8 + PyTorch 2.1。这个组合相对稳定,网上资料也多。PyTorch 2.0之后引入了torch.compile,理论上能加速训练15%-30%,但我在这个项目里没有启用,原因是显存反而会涨一些,而且某些自定义模块编译后会报奇怪的算子错误。如果你求稳,建议先关掉compile,跑通流程再优化。
分词的方案我直接选了OpenAI开源的tiktoken,这是GPT-2原始使用的BPE实现的C版本,速度快,词汇表和GPT-2完全一致。不要自己写BPE训练,因为GPT-2的词汇表是经过特殊处理的,自己训出来的词表会导致模型结构不匹配,且效果大概率变差。
安装的命令很简单:
bash复制pip install torch==2.1.0 tiktoken numpy datasets einops
1.3 整体流程拆解
复现不是一个“写代码然后训练”的简单过程,我把整个项目拆成了六个阶段,每个阶段有明确的产出和验收标准:
| 阶段 | 关键产出 | 验收标准 |
|---|---|---|
| 1. 环境搭建 | 可运行的PyTorch环境 | import torch正常,CUDA可用 |
| 2. 数据准备 | tokenized数据集文件 | 数据格式正确,加载速度达标 |
| 3. 模型实现 | 完整的GPT类 | 前向传播通过,输出形状正确 |
| 4. 训练脚本 | 训练循环 + 日志 | loss稳定下降 |
| 5. 生成脚本 | 采样器 | 能产出可读的文本 |
| 6. 评估调优 | 困惑度指标 | 达到参考水平 |
这个顺序很重要。很多人一上来就写模型,写完才发现数据格式不对,或者训练循环里有个隐性的bug。把每个阶段拆开验收,能大幅减少调试时间。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 数据处理:BPE分词与训练集构建
2.1 BPE分词的核心逻辑
GPT-2使用Byte-Pair Encoding(BPE)分词,词汇表大小50257,其中前256个是单字节(byte)映射,然后是常见的多字符片段,最后是一个特殊的end-of-text标记。为什么需要BPE?最直接的原因是语言模型是在词层面做预测的,但英语单词太多(几十万到几百万),罕见词会导致数据稀疏;而直接按字符建模又会让序列过长,Transformer的注意力计算量随序列长度平方增长,无法接受。
BPE的巧妙之处在于它游走在字符和词之间:高频词直接对应一个token,低频词会被拆成几个子词,完全没见过的词会被拆成字节。举个例子,“unbelievable”可能被拆成“un”、“believ”、“able”三个token,而每个token在训练时都能得到充分的梯度更新。这样既控制了词汇表规模,又保证了模型可以处理任意输入。
2.2 用tiktoken完成分词
tiktoken的用法非常简单:
python复制import tiktoken
enc = tiktoken.get_encoding("gpt2")
text = "Hello, world! This is a test."
tokens = enc.encode(text)
# [15496, 11, 995, 0, 428, 1584, 318, 1110, 13]
decoded = enc.decode(tokens)
# "Hello, world! This is a test."
注意第4个token是0,它代表一个词边界(单词前的空格)。GPT-2的世界里,空格也是token的一部分,这跟很多人的直觉不同。
处理大规模语料时,要按行或按段落增量编码,不要一次性读入整个文件,否则内存会爆。我用的是datasets库从HuggingFace拉取公开数据集,然后并行分词:
python复制from datasets import load_dataset
from tiktoken import get_encoding
from functools import partial
enc = get_encoding("gpt2")
def tokenize_function(examples, enc):
# 将文本拼接后分块,保持最大长度边界
texts = examples["text"]
all_ids = []
for text in texts:
all_ids.extend(enc.encode(text))
all_ids.append(enc.eot_token) # 用<|endoftext|>分隔文档
return {"ids": all_ids}
dataset = load_dataset("openwebtext", split="train", streaming=True)
# dataset = dataset.map(partial(tokenize_function, enc=enc), batched=True)
这里有个隐藏的细节:eot_token(id = 50256)加在每篇文档的末尾。它的作用不只是分隔文档,更关键的是模型会学会在合适的时机输出这个token,相当于学会了“结束一段话”。采样时如果遇到这个token,就要停止生成。
2.3 构建训练批次:连续内存块 vs 随机采样
训练数据有两种组织方式,效果差异很大。第一种是随机抽样:从整个数据集里随机取若干段独立序列,每段长度=1024,不足的部分用padding补零。缺点是会频繁遇到半截子序列和无效的padding,影响训练效率,而且破坏跨文档的上下文连贯性。
第二种是连续内存块:把整个数据集拼成一个超长的token序列,然后从头开始切成连续的块,每块长度1024。这样除了文档边界处,绝大多数样本都是完整的连续文本,信息密度高。我在实现中采用了第二种,代码如下:
python复制import torch
from torch.utils.data import Dataset
class TokenDataset(Dataset):
def __init__(self, tokens, block_size=1024):
self.tokens = tokens
self.block_size = block_size
def __len__(self):
return len(self.tokens) - self.block_size
def __getitem__(self, idx):
chunk = self.tokens[idx: idx + self.block_size + 1]
x = torch.tensor(chunk[:-1], dtype=torch.long)
y = torch.tensor(chunk[1:], dtype=torch.long)
return x, y
输入和标签之间做一个offset:每个位置预测下一个位置的token。这是自回归语言模型的基本范式,整个模型学到的就是在给定前缀的条件下,下一个token的概率分布。
2.4 数据质量与批量加载的工程细节
数据集质量直接决定模型效果。我在实验中确认了几个关键点:清洗时要去掉HTML标签、重复的空行、过短的文本片段;对于中文或非英文内容,要么过滤掉,要么单独处理,因为BPE的词表主要是英文优化的;如果数据集太大,可以先做重复检测,去掉近似重复的文档,能提升训练效率。
DataLoader的num_workers建议设成4-8,prefetch_factor设2,能有效减少GPU等待。如果磁盘是机械硬盘,建议先把tokenized数据转换成.bin格式,用np.memmap来做内存映射读取,避免每次都在数据加载阶段成为瓶颈。
3. 模型架构逐层拆解:从头手写Transformer
3.1 配置类:把超参数集中管理
开始写代码前,先把所有超参数放到一个配置类里。这个习惯在实验阶段特别重要,因为跑一次训练要几个小时,如果超参数散落在代码各处,调参时很容易改错地方。我的配置如下:
python复制from dataclasses import dataclass
@dataclass
class GPTConfig:
vocab_size: int = 50257
n_layer: int = 12
n_head: int = 12
n_embd: int = 768
block_size: int = 1024
dropout: float = 0.1
bias: bool = False
# 训练相关
batch_size: int = 8
learning_rate: float = 3e-4
weight_decay: float = 0.1
warmup_steps: int = 2000
max_steps: int = 60000
这里有个值得注意的点:block_size在GPT-2里也叫n_ctx(上下文长度),它决定了Transformer原始注意力计算时QK^T矩阵的尺寸。在推理时,block_size也固定了模型能处理的最大序列长度,超过这个长度就需要截断或使用更高级的位置编码方案。
3.2 Token嵌入与位置编码
GPT-2只使用了可学习的token嵌入和位置嵌入,两者相加后作为Transformer的输入。和原始Transformer论文不同,它没有使用正弦余弦位置编码,因为可学习的嵌入在数据量足够时能学到相似甚至更好的位置关系。
代码实现:
python复制import torch
import torch.nn as nn
class GPTEmbeddings(nn.Module):
def __init__(self, config):
super().__init__()
self.tok_emb = nn.Embedding(config.vocab_size, config.n_embd)
self.pos_emb = nn.Embedding(config.block_size, config.n_embd)
self.drop = nn.Dropout(config.dropout)
self.block_size = config.block_size
def forward(self, idx):
_, t = idx.shape
assert t <= self.block_size, f"Sequence length {t} exceeds block size {self.block_size}"
pos = torch.arange(0, t, dtype=torch.long, device=idx.device).unsqueeze(0)
tok_emb = self.tok_emb(idx)
pos_emb = self.pos_emb(pos)
return self.drop(tok_emb + pos_emb)
顺带提一句,nn.Embedding本质上是一个查询表,输入token id,输出对应的768维向量。124M参数里,词嵌入占了50257×768≈38.6M,将近三分之一,所以这部分的内存效率也值得关注。
3.3 带因果掩码的多头注意力
Transformer的核心自注意力公式是Attention(Q, K, V) = softmax(QK^T/sqrt(d_k))V。因果掩码的作用是保证每个位置只能看到它之前(包括自己)的token,不能看到未来的信息,这是语言模型自回归性质在训练时的体现。
实现细节里有几个容易出错的地方。第一是缩放因子sqrt(d_k),这里d_k = n_embd / n_head = 768 / 12 = 64,所以缩放因子是8。这个缩放不是拍脑袋定的,它的意义在于控制softmax输入的方差。假设Q和K的元素都是均值为0、方差为1的随机变量,那么QK^T的方差大约是d_k,除以sqrt(d_k)之后方差回到1,softmax不会饱和到梯度消失。第二是mask矩阵的形状,应该是(1, 1, t, t),前面两个1是为了广播到(batch, head)维度。第三是注意力权重上做dropout,这是Transformer原论文里的做法,能起到正则化作用。
python复制class CausalSelfAttention(nn.Module):
def __init__(self, config):
super().__init__()
assert config.n_embd % config.n_head == 0
self.n_head = config.n_head
self.n_embd = config.n_embd
# 合并Q、K、V、输出投影矩阵
self.c_attn = nn.Linear(config.n_embd, 3 * config.n_embd, bias=config.bias)
self.c_proj = nn.Linear(config.n_embd, config.n_embd, bias=config.bias)
self.attn_dropout = nn.Dropout(config.dropout)
self.resid_dropout = nn.Dropout(config.dropout)
self.dropout = config.dropout
def forward(self, x):
B, T, C = x.shape
qkv = self.c_attn(x)
q, k, v = qkv.split(self.n_embd, dim=2)
# 拆成多头,permute让head维度提前
k = k.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
q = q.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
v = v.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
# 注意力分数
att = (q @ k.transpose(-2, -1)) * (1.0 / (k.shape[-1] ** 0.5))
# 因果掩码
causal_mask = torch.tril(torch.ones(T, T, device=x.device, dtype=torch.bool))
att = att.masked_fill(~causal_mask.view(1, 1, T, T), float('-inf'))
att = torch.softmax(att, dim=-1)
att = self.attn_dropout(att)
y = att @ v
y = y.transpose(1, 2).contiguous().view(B, T, C)
return self.resid_dropout(self.c_proj(y))
这里把QKV的线性变换合并到了一个nn.Linear里输出3×768=2304维,而不是分别用三个Linear,主要考虑是矩阵乘法在GPU上是大块运算,合并减少kernel launch次数,训练速度能提升5%-10%。transpose(1,2)之后记得.contiguous(),否则view会报错或者触发非连续内存访问,降低计算效率。
3.4 前馈网络与LayerNorm的实现细节
GPT-2的前馈网络是一个两层的MLP:先把维度从768放大到4×768=3072,用GELU激活,再压缩回768。这个扩展比例4倍在GPT系列中是固定的,后来很多模型也沿用这个思路。
GELU(Gaussian Error Linear Unit)是一个平滑版的ReLU,公式是x * Phi(x),其中Phi是标准正态分布的CDF。实现时用tanh近似:
python复制class MLP(nn.Module):
def __init__(self, config):
super().__init__()
self.c_fc = nn.Linear(config.n_embd, 4 * config.n_embd, bias=config.bias)
self.c_proj = nn.Linear(4 * config.n_embd, config.n_embd, bias=config.bias)
self.dropout = nn.Dropout(config.dropout)
def forward(self, x):
x = self.c_fc(x)
x = torch.nn.functional.gelu(x, approximate='tanh')
return self.dropout(self.c_proj(x))
LayerNorm放在注意力层和MLP之前,也就是GPT-2使用“Pre-LN”结构:x = x + sublayer(layernorm(x))。这和原始Transformer的Post-LN相反。Pre-LN的优点是训练更稳定,对大学习率不敏感,但下限略低;Post-LN如果超参数没调好,容易出现梯度爆炸或消失。现在的模型基本全用Pre-LN,LLaMA也是这个结构。
还要补充一点:nn.LayerNorm默认的bias=True,但现在的GPT-2复现一般不加bias,因为后续的激活函数已经提供了非线性;去掉bias能稍微节省显存,也能让模型更依赖layer norm的均值和方差归一化能力。
3.5 组装TransformerBlock与完整GPT模型
TransformerBlock就是“多头注意力 + 残差”拼接“前馈 + 残差”:
python复制class Block(nn.Module):
def __init__(self, config):
super().__init__()
self.ln_1 = nn.LayerNorm(config.n_embd)
self.attn = CausalSelfAttention(config)
self.ln_2 = nn.LayerNorm(config.n_embd)
self.mlp = MLP(config)
def forward(self, x):
x = x + self.attn(self.ln_1(x))
x = x + self.mlp(self.ln_2(x))
return x
整个GPT模型在最前面接嵌入层,中间是12个Block的堆叠,最后接一个LayerNorm和一个线性层把768维映射到50257个词表项上。最后的线性层和词嵌入矩阵如果做权重共享(weight tying),参数量能省下38.6M,从约163M降到124M左右,这正是“124M”这个数字的来源。GPT-2官方是共享的,实际效果也验证了这一点:共享之后模型在少一个矩阵参数的情况下表现没有明显下降,甚至在小样本上还略有正则化效果。
python复制class GPT(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
self.transformer = nn.ModuleDict(dict(
wte=nn.Embedding(config.vocab_size, config.n_embd),
wpe=nn.Embedding(config.block_size, config.n_embd),
drop=nn.Dropout(config.dropout),
h=nn.ModuleList([Block(config) for _ in range(config.n_layer)]),
ln_f=nn.LayerNorm(config.n_embd),
))
self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)
self.lm_head.weight = self.transformer.wte.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)
def forward(self, idx):
B, T = idx.shape
x = self.transformer.wte(idx) + self.transformer.wpe(torch.arange(T, device=idx.device))
x = self.transformer.drop(x)
for block in self.transformer.h:
x = block(x)
x = self.transformer.ln_f(x)
logits = self.lm_head(x)
return logits
参数初始化用了标准差0.02的正态分布,另外还有一个经验做法是残差分支的最后一层(c_proj)用标准差0.02/sqrt(2*n_layer)来初始化,代码里没有显式区分,但模块的默认初始化已经够用了。如果训练不稳定,可以尝试加上这个更精细的初始化。
3.6 前向传播与Loss计算
Loss计算是下一步:logits形状是(B, T, vocab_size),标签是(B, T)。PyTorch的nn.CrossEntropyLoss接受任意维度的输入和标签,但需要把logits的前两维合并:
python复制def compute_loss(model, x, y):
logits = model(x) # (B, T, vocab_size)
loss = nn.functional.cross_entropy(
logits.view(-1, logits.size(-1)),
y.view(-1)
)
return loss
这里每个位置都进行了一次预测,相当于B×T个独立的分类任务,每个任务的目标都是预测下一个token。对一个真实的1024长度序列来说,这意味着模型在1024个位置上同时学习预测,一个batch的数据量带来的梯度更新效力相当于单任务训练的1024倍,这也是Transformer语言模型数据利用效率高的重要原因。
4. 训练循环与优化策略
4.1 Warmup与Cosine学习率调度
Language model训练中,学习率调度可能是对收敛效果影响最大的工程细节。124M模型的实践配置是:前2000步线性warmup,从0逐步升到3e-4,然后用余弦退火降到接近0。
为什么不从一开始就用大学习率?如果学习率过高,初始阶段参数的剧烈波动会让模型在一个很差的区域里徘徊,之后很难恢复正常。Warmup让模型先用小步幅“试探”一下loss landscape的粗糙程度,稳定后再加速。为什么不一直保持一个固定学习率?因为在训练后期,参数已经接近一个较优的区域,过大的步长会在最优解附近反复震荡,无法精确收敛。余弦退火在前期保持较高学习率加速收敛,后期逐渐衰减帮助模型进到更细的谷底。
python复制import math
def get_lr(it, config):
if it < config.warmup_steps:
return config.learning_rate * (it + 1) / config.warmup_steps
if it > config.max_steps:
return 0
progress = (it - config.warmup_steps) / (config.max_steps - config.warmup_steps)
return 0.5 * config.learning_rate * (1 + math.cos(math.pi * progress))
4.2 优化器:AdamW与权重衰减分组
AdamW是当前Transformer训练的标配。相比Adam,它把权重衰减从梯度动量里分离开,单独对参数做L2惩罚,避免了Adam里L2正则化被二阶动量缩放带来的耦合效应。权重衰减系数设为0.1,这个值远大于很多传统模型使用的1e-4或1e-2,因为现代大模型普遍认为权重衰减在这里更多是扮演“稳定性保险”的角色,对泛化能力的贡献是次要的。
工程上的关键是参数分组:只对2D以上的矩阵参数(Embedding和Linear的weight)做权重衰减,不对bias和LayerNorm的gamma、beta做。原因是这些一维参数本身量级很小,权重衰减反而会削弱它们的表达能力。实现:
python复制def configure_optimizer(model, weight_decay, learning_rate, betas=(0.9, 0.95)):
decay_params = []
no_decay_params = []
for name, param in model.named_parameters():
if param.ndim >= 2:
decay_params.append(param)
else:
no_decay_params.append(param)
optimizer = torch.optim.AdamW([
{"params": decay_params, "weight_decay": weight_decay},
{"params": no_decay_params, "weight_decay": 0.0},
], lr=learning_rate, betas=betas)
return optimizer
beta2设成0.95而不是默认的0.999,这个选择也不是随便定的。训练Transformer时,梯度并不总是平稳的,有时会出现大的梯度尖峰,如果beta2太大,二阶动量的历史估计会拖慢对尖峰的响应,导致训练不稳定。0.95能让优化器更快适应梯度的变化。
4.3 混合精度与梯度累积
在24GB显存上跑batch size 8×1024已经是极限,但更大的batch通常能带来更稳定的梯度。一个折中方案是梯度累积:每grad_accum_steps步把梯度累加起来,再统一更新参数,模拟更大的batch size。我实际配置了grad_accum_steps=4,相当于有效batch size为32。
混合精度(AMP)是另一大显存节省手段。PyTorch的自动混合精度(Automatic Mixed Precision)API非常方便,我用的策略是torch.cuda.amp.GradScaler配合autocast。核心思路:前向和反向在float16下计算,大幅减少显存和算力消耗;梯度更新时通过scaler缩放,避免small gradient在下溢为0;参数更新前再把梯度转回float32计算。float16的表示范围是±65504,如果梯度里出现很小的值(比如1e-5),在float16里就直接变成0了,所以用scaler放大损失梯度,更新前再缩小。
python复制scaler = torch.cuda.amp.GradScaler()
for step in range(max_steps):
lr = get_lr(step, config)
for param_group in optimizer.param_groups:
param_group['lr'] = lr
for _ in range(grad_accum_steps):
x, y = next(train_iter)
with torch.cuda.amp.autocast():
logits = model(x)
loss = torch.nn.functional.cross_entropy(
logits.view(-1, logits.size(-1)),
y.view(-1)
) / grad_accum_steps # 累积步数上的平均
scaler.scale(loss).backward()
# 梯度裁剪,防止梯度爆炸
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
if step % 100 == 0:
print(f"step {step}, loss {loss.item() * grad_accum_steps:.4f}, lr {lr:.2e}")
梯度裁剪这里踩过一个坑:如果用scaler.scale(loss).backward(),梯度已经是缩放后的值,直接做clip_grad_norm_会出错。必须先scaler.unscale_(optimizer)把梯度还原为float32,再做裁剪。顺序错了,梯度裁剪就完全失效了,甚至可能因为float16下溢出导致NaN。
4.4 训练日志与检查点管理
训练日志要记录至少这些信息:step、loss、learning_rate、GPU显存占用、每个token的平均耗时。用简单的Python字典打印就行,不需要引入tensorboard这类重量级工具。检查点保存我用了两种策略:每1000步保存一个最新的checkpoint,训练结束保存一个final checkpoint。保存时用torch.save把模型参数、优化器状态、step、loss都存下来,这样即使训练中断也能恢复:
python复制checkpoint = {
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'scaler_state_dict': scaler.state_dict(),
'step': step,
'config': config,
}
torch.save(checkpoint, f'checkpoints/gpt2_124m_step{step}.pt')
恢复训练的代码同样要处理scaler的状态,否则混合精度的缩放因子信息丢失后,后续的scale值会重新从65536开始,如果当前梯度范围已经不同,会导致前几步的梯度干脆变成0。
4.5 多卡分布式训练的可选方案
如果你的机器有多张GPU,可以尝试torch.nn.DistributedDataParallel(DDP)。124M模型在多卡训练时,通信开销占总训练时间的比例在2张卡上非常小,4张卡开始有明显加速。DDP的核心是每个GPU持有模型的一个完整副本,但数据是分片的,每个GPU维护自己的前向和反向计算,反向传播完会通过环状通信(Ring All-Reduce)同步梯度,然后各自更新参数。
实现的额外步骤包括:
python复制# 初始化
torch.distributed.init_process_group(backend='nccl')
torch.cuda.set_device(local_rank)
model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[local_rank])
# DataLoader需要设置
# sampler = torch.utils.data.distributed.DistributedSampler(dataset)
一个常见的误区是数据分片时没有将每个batch均匀分配给不同GPU,导致负载不均衡。用DistributedSampler能确保每个epoch里每个GPU拿到的样本不重叠,且覆盖整个数据集。
5. 训练过程中的问题排查与性能调优
5.1 Loss不下降的排查顺序
我的loss曲线在最开始也遇到过“死活不降”的情况。排除顺序建议是:先确认损失函数本身是否在合理范围内。对于GPT-2 124M词汇表50257,随机初始化的模型在第一个step的loss应该约等于log(50257)≈10.82。如果初始loss远大于这个值,说明输出层初始化有问题;如果初始loss就小于10,很可能标签和数据不对齐,模型在记忆某种捷径。
然后是检查数据流:随机抽几个batch打印输入token和对应的output token,人工核对一下是不是“预测下一个token”的逻辑。接着看梯度:如果梯度范数很大(超过几个数量级),可能是残差结构、LayerNorm的位置或初始化有问题;如果梯度范数接近0,则可能模型某些部分被dropout完全关闭了,或者Float16下梯度下溢。最后才是学习率的问题,把warmup和lr打出来,确认实际生效的学习率确实在上升。
5.2 显存不足的应急方案
24GB显存跑batch size 8已经比较满,如果遇到OOM(Out of Memory),最有效的调整是依次尝试:降低batch size到4或2;开启torch.utils.checkpoint.checkpoint(梯度检查点),用重计算换显存;把block_size从1024降到512;换成更小的float16(如果用的是fp32训练)。梯度检查点的原理是在反向传播时不保存前向的激活值,而是在反向时重新算一遍,显存占用大约从O(L)降低到O(sqrt(L)),但训练时间会增加30%-50%,属于保命方案。
显存不够时还有一种“投机取巧”的办法:把batch size降到1,配合梯度累积到32步。这样有效batch size仍然可以保持在32,只是每个batch的样本数少了,梯度噪声会大一些,但训练还是能稳定进行的。
5.3 训练崩溃与NaN的处理
NaN的根源通常是损失中出现无穷值。排查链路的顺序是:
- 混合精度:如果开了AMP,先关掉跑几步,确认是不是float16的数值稳定性问题。float16最大的坑在于极端情况下的梯度下溢或溢出,解决办法是加大缩放因子,或在LayerNorm的epsilon上做文章(从1e-5改成1e-6)。
- 学习率:如果l r设置得太高,权重更新一次就可能破坏数值稳定性。赶紧看日志里lr的数值,以及loss在NaN之前的走势。
- 梯度裁剪:如果某个batch的梯度出现了异常的分量,梯度裁剪能挡住大部分NaN的来源。把max_norm从1.0改成0.5或0.3。
- 数据本身:某些输入里包含了不合法的token(比如超出词汇表范围),会让Embedding查询越界,产生NaN。在数据集处理时做一个token id的边界检查。
有些NaN是“间歇性”的,比如连续跑了2000步没事,突然一个特殊的batch触发崩溃。这时候最朴素也有效的办法是:找到崩溃对应的batch,单独跑一遍前向反向,逐步把batch缩小,定位出具体是哪个样本触发的,然后检查那个样本的数据。
5.4 性能瓶颈与吞吐量优化
训练速度上,我做了几个实验对比。在单张RTX 3090上,batch size 8、序列长度1024、混合精度,模型的理论算力利用率大约只有30%-40%。瓶颈主要在数据加载和attention计算。提升手段包括:
- 数据加载:用
torch.utils.data.DataLoader的persistent_workers=True,避免每次epoch重复启动worker进程;数据集先缓存到内存(或者内存映射),避免磁盘IO成为瓶颈。 - attention计算:在PyTorch 2.1里,把attn换成
torch.nn.functional.scaled_dot_product_attention,能利用FlashAttention的底层优化,训练速度大约提升20%-30%,显存还能再省一点。这个方法我在项目后半段用上了,改动量非常小。 - 使用
torch.backends.cudnn.benchmark = True:对固定尺寸的模型可以让cuDNN自动搜索最合适的卷积/矩阵乘法kernel。
FlashAttention的原理是分块计算注意力矩阵,避免显式构建(B, T, T)的完整注意力矩阵。对1024长度来说,一个head的注意力矩阵是1024×1024,12个head × B个batch×12层,累计下来是一个很大的显存占用。FlashAttention把这个矩阵分块后在SRAM里算完,磁盘都不落,省去了大量HBM读写。
6. 生成采样与评估指标
6.1 自回归生成原理
训练完之后,生成是一个自回归的过程:给模型一个初始的prompt,它输出下一个token的概率分布,然后我们采样一个token,把它拼到prompt末尾,再输入模型得到下一个token。重复这个过程,直到达到最大长度或遇到<|endoftext|>标记。
核心采样配置是top-k和top-p(nucleus sampling)。top-k的思路是只从概率最高的k个token里按概率采样,k通常取50;top-p是选择累计概率达到p的最小token集合,p通常取0.9-0.95。两者都使用而不是只用temperature,是因为temperature会改变整体概率分布的锐利程度,但可能会让低概率的长尾token依然有机会被采样,而top-k/top-p直接截断了分布在尾部的那批非常不合理的token。
python复制def generate(model, prompt, max_new_tokens=100, temperature=0.8, top_k=50, top_p=0.95):
model.eval()
enc = tiktoken.get_encoding("gpt2")
tokens = enc.encode(prompt)
input_ids = torch.tensor([tokens], dtype=torch.long, device='cuda')
with torch.no_grad():
for _ in range(max_new_tokens):
# 只取最后block_size个token,防止超长
input_ids = input_ids[:, -config.block_size:]
logits = model(input_ids)
next_logits = logits[0, -1, :] / temperature
if top_k > 0:
v, _ = torch.topk(next_logits, min(top_k, next_logits.size(-1)))
next_logits[next_logits < v[-1]] = -float('Inf')
if top_p > 0:
sorted_logits, sorted_indices = torch.sort(next_logits, descending=True)
cumulative_probs = torch.cumsum(torch.nn.functional.softmax(sorted_logits, dim=-1), dim=-1)
sorted_mask = cumulative_probs > top_p
sorted_mask[..., 1:] = sorted_mask[..., :-1].clone()
sorted_mask[..., 0] = False
sorted_logits[sorted_mask] = -float('Inf')
next_logits = torch.zeros_like(next_logits).scatter_(0, sorted_indices, sorted_logits)
probs = torch.nn.functional.softmax(next_logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
input_ids = torch.cat([input_ids, next_token], dim=-1)
output_tokens = input_ids[0].tolist()
return enc.decode(output_tokens)
6.2 快速验证训练效果:从“乱码”到“通顺”
如果训练充分,模型生成的文本应该具备基本的英语语法和连贯的语义。如果只训练了10000步,生成结果往往是几个词的重复循环,这不是bug,说明模型学到的上下文是短期的;继续训练到50000步以上,输出会明显变长而且语义相对连贯。
要快速验证模型是否学到了真正的语言结构(而不是死记硬背),可以做几组测试:
- 长度泛化:给模型一个很短的prompt(比如“Hello”),看能否生成一段完整的话。
- 输入变化:换不同的prompt风格,比如问句、陈述句、新闻报道的开头,模型是否相应调整了输出的语体。
- 扰动测试:给prompt加一个打字错误,模型是否还能正常继续。如果模型过度依赖拼写而缺乏语义理解,可能乱码输出。
另外,用同一个prompt跑多次生成,观察输出多样性。如果每次输出都几乎一样,说明temperature设置得太低或者模型的概率分布过于突出某个token;如果输出之间完全没有重叠,说明temperature太高,模型可能在随机游走。
6.3 评估指标:困惑度(Perplexity)
困惑度是语言模型最常用的评估指标,它直接和交叉熵loss挂钩:PPL = exp(loss)。在训练集上,模型收敛的loss大约对应PPL在10-15之间;在验证集上,如果PPL超过训练集的1.5倍,说明可能过拟合了。
困惑度的一个直观解释是:PPL等于模型在每个位置上平均“不知所措”的选择数。如果PPL=20,说明模型在每一步预测下一个token时,相当于从20个候选里随机猜一个。PPL越低越好,但也不是越低越好,过低的PPL可能意味着模型只是在死记训练文本的分布模式,缺少泛化能力。
计算验证PPL时要注意:把验证集切成长度一致的块(不跨越文档边界,eot token处要切断),然后用同训练一致的batch组织方式跑一次前向,累计所有位置的loss,最后取平均再exp。这个过程只做前向,不做反向,也不更新权重。
6.4 与HuggingFace预训练模型做对比
一个有趣的基准测试是:把训练出来的模型和HuggingFace上的gpt2(相同124M规模)在几个固定prompt上做输出对比。我们的模型在训练数据量远小于原始GPT-2(原始GPT-2用了约40GB的WebText,我们用了几GB的公开数据)的情况下,PPL会有差距,但生成文本的句法和词汇多样性应该已经能达到“半专业”水准。
这种差距是正常的,不说明复现失败。通过这个对比,还能从数据量、数据质量、训练步数等维度定位出提升空间,这本身就是复现项目最大的收获。
7. 复现成果与后续扩展方向
7.1 这次复现的实际数据
我在单张RTX 3090上跑完了60000步,总耗时约9天(中间因为显存问题和一次断电中断过两次,用了checkpoint恢复)。最终验证集的交叉熵loss约2.9,对应PPL约18.2,生成效果已经能生成语法基本正确的英文短文本,但长文本的连贯性还与OpenAI原始权重有明显差距。
| 指标 | 我的复现 | OpenAI GPT-2(参考) |
|---|---|---|
| 参数总量 | 124M | 124M |
| 训练数据 | 约8GB公开文本 | 约40GB WebText |
| 训练步数 | 60000 | 更高(配置不同) |
| 验证PPL | 18.2 | 约14-15 |
| 单段生成最大长度 | 1024 token | 1024 token |
需要说明的是,PPL的差异主要来自训练数据规模和训练预算。如果训练数据提升到30GB以上,PPL会明显下降。这也印证了大模型训练的铁律:数据量和模型容量同等重要。
7.2 训练技巧总结
实践中最重要的五个技巧,按有效程度排序:
- 权重共享(weight tying)省了38M参数,而效果几乎没有下降,这是最划算的一笔优化。
- 混合精度训练(AMP)让显存需求几乎减半,同时训练速度提升约40%。如果不用AMP,24GB显存跑batch size 8很紧张,要在小batch和长序列之间反复挣扎。
- 连续内存块式数据集比随机采样训练效率高得多,且实现简单。
- Warmup + Cosine学习率调度是稳定训练的基石,其重要性远超模型结构的微调。
- 梯度裁剪保住了掉进深坑的模型好几次;这也是AdamW与clip组合使用的最佳实践。
7.3 从124M继续扩展的路线
跑通124M之后,往两个方向走都有价值。一个方向是横向对比:把n_layer改成6或24,n_embd改成384或1024,观察不同规模的学习曲线差异,这能直观感受“参数增加带来的边际收益递减”。另一个方向是纵向改进:加入FlashAttention、使用旋转位置编码(RoPE)、把LayerNorm换成RMSNorm,等于做一次“从GPT-2到LLaMA架构演进”的手工课。
如果想把模型用在具体任务上,可以在复现基础上做监督微调(SFT):准备一批指令-回答对,在这个基座模型上接着训练几天,模型就能学会“对话体”的生成风格。这一步本质上就是当年ChatGPT训练链路里的“第一步”。
7.4 如果想在更小显存上开始
如果你的显卡只有8GB或12GB显存,也不用灰心。把batch size降到2、block_size设512、关闭部分dropout,124M模型在8GB上也能跑起来,只是训练时间会拉长不少。如果8GB仍然OOM,可以考虑把n_layer减半到6层(约60M参数),这仍然保留了GPT-2的核心结构,作为学习的载体完全足够。
7.5 自己复现和直接调用库的体验差异
我自己在做这个项目的过程中,最大的体会是:用HuggingFace的from_pretrained加载一个GPT-2只要几秒钟,但自己动手“搭积木”一遍,才真正理解了什么叫做“所有魔法背后都是矩阵乘法”。现在再看网上关于大模型的讨论,很多细节都能对上了:为什么显存不够要梯度检查点、为什么训练要用warmup、为什么推理时要做top-p采样。如果你也是“用了很久Transformer但觉得自己只是API调用者”,我强烈建议花两周按这个路线复现一次,收获绝对超出预期。
最后再分享一个小建议:复现过程一定要养成保存checkpoint和记录实验日志的习惯。我中途有一次训练崩了,当时幸好checkpoint保存在2000步前,否则那几天的算力就白白浪费了。训练类的项目,代码的稳定性和可恢复性永远比“把代码写得短”重要得多。
