从零开始训练你自己的大语言模型:完整端到端指南
一份动手教程,教你使用PyTorch从零构建、预训练并对齐Transformer大语言模型,涵盖SFT、DPO、PPO和GRPO。
你是否曾想从零开始训练自己的大语言模型(LLM),而不仅仅是微调现有模型?FareedKhan-dev 的开源项目 train-llm-from-scratch 提供了一个完整、透明的流水线,带你从原始文本到对齐的、具备推理能力的模型。这不是一个围绕现有库的黑盒封装,而是一次动手教育之旅,其中的每个算法——从多头注意力到GRPO——都用纯PyTorch实现。
本指南将带你走完整个过程,解释每一步的“为什么”和“如何”,让你理解基本原理并亲自运行代码。
为什么要从零构建LLM?
大多数开发者通过Hugging Face的 transformers 或 trl 等高层次API与LLM交互。虽然这些工具功能强大,但它们抽象了核心机制。从零构建能让你:
- 深入理解: 你确切了解分词、注意力和训练循环是如何工作的。
- 完全控制: 你可以修改任何组件——损失函数、架构或训练策略——而不受库的限制。
- 教育价值: 这是理解因果掩码、奖励模型和策略梯度方法的最佳途径。
该项目专为希望在一个地方看到完整流水线(从数据下载到文本生成)的学生、开发者和研究人员设计。
流水线概览
整个流程遵循清晰的顺序路径:
- 数据准备: 原始文本 → 令牌ID → 存储数组。
- 模型构建: 从小型可复用组件(MLP、注意力、块)构建Transformer。
- 预训练: 在大语料库(The Pile)上进行下一个令牌预测。
- 文本生成: 从训练好的模型中采样。
- 后训练(对齐): 通过SFT、奖励建模、DPO、PPO和GRPO将基础模型转变为有用的助手。
- 评估: 在GSM8K等基准上测量性能。
- 推理与聊天: 与最终模型交互。
第1步:准备数据
模型只理解整数。第一步是将文本转换为令牌ID并高效存储。
分词
该项目使用OpenAI的 r50k_base 分词器(来自 tiktoken),与GPT-3使用的相同。文本变成整数列表,并在每个文档后附加一个特殊的 <|endoftext|> 令牌(ID 50256),以便模型学习边界。
# 从The Pile下载并分词数据
python scripts/data_download.py
python scripts/data_preprocess.py
对于更新、更快的路径:
python scripts/prepare_pretrain_data.py --split val --out data/pile_dev.h5
python scripts/prepare_pretrain_data.py --split train --num_shards 1 --out data/pile_train.h5
聊天格式与损失掩码
对于后训练,模型需要理解角色(用户 vs. 助手)。该项目使用纯文本标记:
<|user|>
13 + 29 等于多少?<|endoftext|><|assistant|>
<think>13 + 29 = 42</think><answer>42</answer><|endoftext|>
损失掩码确保模型只从助手的令牌学习,而不从用户的提示学习。这对SFT和RL至关重要。
def encode_chat(messages, add_generation_prompt=False):
ids, mask = [], []
for m in messages:
role = m["role"]
header_ids = _encode_ordinary(_header_for(role))
ids.extend(header_ids)
mask.extend([0] * len(header_ids))
content_ids = _encode_ordinary(m["content"])
is_completion = role == "assistant"
ids.extend(content_ids)
mask.extend([1 if is_completion else 0] * len(content_ids))
ids.append(EOT_ID)
mask.append(1 if is_completion else 0)
return ids, mask
第2步:构建Transformer
Transformer由四个小型、可复用的PyTorch模块构建而成。这种模块化设计使代码易于理解和修改。
多层感知机(MLP)
每个块的“思考”部分。它扩展令牌向量,应用ReLU,然后投影回原尺寸。
class MLP(nn.Module):
def __init__(self, n_embed):
super().__init__()
self.hidden = nn.Linear(n_embed, 4 * n_embed)
self.relu = nn.ReLU()
self.proj = nn.Linear(4 * n_embed, n_embed)
def forward(self, x):
x = self.relu(self.hidden(x))
x = self.proj(x)
return x
单头注意力
允许一个令牌查看其他令牌。它计算查询、键和值向量,应用因果掩码(使令牌不能看到未来令牌),并取值的加权和。
class Head(nn.Module):
def __init__(self, head_size, n_embed, context_length):
super().__init__()
self.key = nn.Linear(n_embed, head_size, bias=False)
self.query = nn.Linear(n_embed, head_size, bias=False)
self.value = nn.Linear(n_embed, head_size, bias=False)
self.register_buffer('tril', torch.tril(torch.ones(context_length, context_length)))
def forward(self, x):
B, T, C = x.shape
k = self.key(x)
q = self.query(x)
scale_factor = 1 / math.sqrt(C)
attn_weights = q @ k.transpose(-2, -1) * scale_factor
attn_weights = attn_weights.masked_fill(self.tril[:T, :T] == 0, float('-inf'))
attn_weights = F.softmax(attn_weights, dim=-1)
v = self.value(x)
out = attn_weights @ v
return out
多头注意力
并行运行多个注意力头,每个头学习不同的关系,然后拼接并投影结果。
Transformer块
结合注意力和MLP,并带有预归一化残差连接。这是核心重复单元。
class Block(nn.Module):
def __init__(self, n_head, n_embed, context_length):
super().__init__()
self.ln1 = nn.LayerNorm(n_embed)
self.attn = MultiHeadAttention(n_head, n_embed, context_length)
self.ln2 = nn.LayerNorm(n_embed)
self.mlp = MLP(n_embed)
def forward(self, x):
x = x + self.attn(self.ln1(x))
x = x + self.mlp(self.ln2(x))
return x
完整Transformer
封装所有内容:令牌嵌入、位置嵌入、块堆栈、最终层归一化以及到词汇表大小logits的线性投影。
class Transformer(nn.Module):
def __init__(self, n_head, n_embed, context_length, vocab_size, N_BLOCKS):
super().__init__()
self.token_embed = nn.Embedding(vocab_size, n_embed)
self.position_embed = nn.Embedding(context_length, n_embed)
self.attn_blocks = nn.ModuleList([Block(n_head, n_embed, context_length) for _ in range(N_BLOCKS)])
self.layer_norm = nn.LayerNorm(n_embed)
self.lm_head = nn.Linear(n_embed, vocab_size)
self.register_buffer('pos_idxs', torch.arange(context_length))
def forward(self, idx, targets=None):
x = self.forward_hidden(idx)
logits = self.lm_head(x)
loss = None
if targets is not None:
B, T, C = logits.shape
flat_logits = logits.reshape(B * T, C)
targets = targets.reshape(B * T).long()
loss = F.cross_entropy(flat_logits, targets)
return logits, loss
第3步:预训练基础模型
预训练是在大型语料库上进行下一个令牌预测。模型读取随机窗口的令牌,预测每个位置的下一个令牌,并通过交叉熵损失更新权重。
# 训练一个1300万参数的模型
python scripts/train_transformer.py
# 对于更大的模型,使用内存优化
python scripts/train_transformer.py --amp --grad-checkpointing --grad-accum 8
损失从接近 ln(vocab_size)(约10.8)开始,随着模型学习语言统计特性而下降。一个7700万参数的模型在2块L40 GPU上训练,经过2000步后损失达到约3.73。
第4步:生成文本
训练完成后,你可以通过从模型的输出分布中采样来生成文本。
def generate(self, idx, max_new_tokens):
for _ in range(max_new_tokens):
idx_cond = idx[:, -self.context_length:]
logits, _ = self(idx_cond)
logits = logits[:, -1, :]
probs = F.softmax(logits, dim=-1)
idx_next = torch.multinomial(probs, num_samples=1)
idx = torch.cat((idx, idx_next), dim=1)
return idx
python scripts/generate_text.py --model_path models/transformer_B.pt --input_text "The" --max_new_tokens 100
第5步:后训练(对齐)
基础模型可以继续文本,但不能遵循指令。后训练通过监督微调(SFT)和强化学习(RL)使其对齐。
SFT(监督微调)
教模型以聊天格式回答。损失仅在助手的令牌上计算。
python scripts/prepare_sft_data.py --context_length 1024
torchrun --standalone --nproc_per_node=2 scripts/train_sft.py
奖励模型
在SFT模型之上训练一个小型线性头来评分答案。使用Bradley-Terry损失来偏好选择的回答而非拒绝的回答。
def bradley_terry_loss(chosen_rewards, rejected_rewards):
return -F.logsigmoid(chosen_rewards - rejected_rewards).mean()
DPO、ORPO和KTO
直接偏好优化(DPO)跳过奖励模型和RL循环,直接在偏好对上工作。ORPO和KTO是其变体。
torchrun --standalone --nproc_per_node=2 scripts/train_dpo.py --loss_type dpo
PPO(近端策略优化)
经典的RLHF循环。模型生成答案,用奖励模型评分,并通过裁剪的策略损失进行更新。
def ppo_policy_loss(new_logp, old_logp, advantages, mask, clip=0.2):
ratio = torch.exp(new_logp - old_logp)
surr1 = ratio * advantages
surr2 = torch.clamp(ratio, 1.0 - clip, 1.0 + clip) * advantages
loss = -masked_mean(torch.min(surr1, surr2), mask)
return loss
GRPO(组相对策略优化)
DeepSeek-R1风格的方法。它为每个提示采样一组答案,进行评分,并使用组的均值和标准差作为基线。
def group_advantages(rewards, group_size, eps=1e-4):
r = rewards.view(-1, group_size)
mean = r.mean(dim=1, keepdim=True)
std = r.std(dim=1, keepdim=True)
adv = (r - mean) / (std + eps)
return adv.reshape(-1)
第6步:评估
关键指标是贪婪GSM8K准确率:模型必须在 <answer> 标签内输出正确答案。
for s in base_pretrained sft dpo ppo grpo; do
python scripts/eval_post_training.py --ckpt models/$s.pt --label $s --limit 200 --append logs/table.jsonl
done
第7步:与模型对话
使用聊天脚本与任何检查点交互。
python scripts/chat.py --ckpt models/sft.pt --prompt "13 + 29 等于多少?"
开始使用
克隆仓库:
git clone https://github.com/FareedKhan-dev/train-llm-from-scratch.git cd train-llm-from-scratch pip install -e .从小处着手: 首先训练1300万参数的模型。它很快,能让你感受整个流水线。
扩展规模: 增加
n_embed和n_blocks,直到达到GPU内存限制。走完后训练链: 依次运行SFT、DPO、PPO或GRPO,观察GSM8K准确率提升。
这个项目是一颗罕见的宝石:它既是实用工具,又是全面教程。无论你是想深入理解Transformer的学生,还是希望尝试对齐技术的开发者,train-llm-from-scratch 都是完美的起点。
来源
FareedKhan-dev/train-llm-from-scratch: 一种训练你的LLM的直截了当的方法,从下载数据到生成文本。