fareedkhan-dev/train-llm-from-scratch · 上手攻略
- 仓库:fareedkhan-dev/train-llm-from-scratch
- 链接:https://github.com/FareedKhan-dev/train-llm-from-scratch
- 分类:ai
- 作者:Tom
- 更新:2026-07-10
这是什么
一个从零开始训练完整大语言模型(LLM)的全链路教程仓库。手写 PyTorch 实现,不依赖 transformers、trl、peft 等高层库,从文本语料处理讲起,一路走到 SFT(监督微调)→ 奖励模型 → RLHF(PPO、DPO、GRPO),覆盖一个小型 LLM 从诞生到"会聊天"的全过程。
项目作者 FareedKhan-dev 目前在寻找 AI 方向的 PhD 位置,整个项目有着很强的教学感:每段代码前都有文字解释,每个阶段都有真实运行的输出供参考。
解决什么问题
市面上的 LLM 训练教程往往依赖高层库,隐藏了大量实现细节。这个仓库的目标是:让读者真正理解 LLM 训练的每一个环节,从最底层的 tokenization、Attention 实现,到预训练的 loss mask,再到 RLHF 的数学直觉,全部手写展示。
另一个痛点是:大多数教程只讲"怎么跑通",不讲"为什么会这样"。这个仓库用颜色编码图示(绿=原始数据、青=存储数据、蓝=处理步骤、黄=模型/训练、橙=RL 奖励、红=loss、紫=输出/评估)让整个 pipeline 一目了然。
快速安装
git clone https://github.com/FareedKhan-dev/train-llm-from-scratch.git
cd train-llm-from-scratch
pip install -e .
可选组件:
pip install -e ".[train]" # datasets + wandb(日志)
pip install -e ".[ui]" # streamlit 控制面板
pip install -e ".[docs]" # mkdocs 文档站
pip install -e ".[all]" # 全部安装
硬件要求(来自仓库 README,直接可用 T4 GPU 训练 13M 参数模型):
| GPU | 显存 | 2B 训练 | 13M 训练 | 最大实用参数 |
|---|---|---|---|---|
| NVIDIA A100 | 40 GB | ✅ | ✅ | ~6B–8B |
| NVIDIA RTX 4090 | 24 GB | ✅ | ✅ | ~4B |
| NVIDIA RTX 3090 | 24 GB | ✅ | ✅ | ~3.5B–4B |
| NVIDIA RTX 4080 | 16 GB | ❌ | ✅ | ~2B |
| Tesla T4(Colab 免费) | 16 GB | ❌ | ✅ | ~1.5B–2B |
遇到 OOM 时可加 --amp(混合精度)、--grad-checkpointing(梯度检查点)、--grad-accum(梯度累积)降低显存占用。
核心用法
1. 数据准备
预训练数据(The Pile 子集):
# 下载 + tokenize(Legacy 路径)
python scripts/data_download.py
python scripts/data_preprocess.py
# 输出: data/train/pile_train.h5, data/val/pile_dev.h5
# 更快的流式路径
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
SFT / RLHF 数据准备:
python scripts/prepare_sft_data.py # Alpaca + Dolly + GSM8K → sft_packed.h5
python scripts/prepare_preference_data.py # HH-RLHF + UltraFeedback → preferences.jsonl
python scripts/prepare_rl_prompts.py # GSM8K + 算术 → rl_prompts.jsonl
2. 预训练
# Legacy 小规模(13M 参数,smoke 测试)
python scripts/train_transformer.py
# 新路径(可配置更大模型)
# 编辑 configs/ 下的 JSON 文件,如 configs/pretrain/base.json
python scripts/train.py --config configs/pretrain/base.json --lr 1e-3 --batch_size 8
# 快速冒烟测试(极小配置,CPU 也能跑完)
python scripts/train.py --config configs/smoke/pretrain_smoke.json
3. 文本生成
from src.post_training.inference import generate
text = generate("The capital of France is", model_path="checkpoints/pretrain/model.pt")
print(text)
4. SFT(监督微调)
核心思想:对齐(alignment)时只训练 assistant 回复部分,用 0/1 mask 遮住 system prompt 和 user prompt。关键代码逻辑(来自仓库 README):
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)) # header 不参与训练
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
运行 SFT:
python scripts/sft.py --config configs/sft/base.json --lr 2e-5 --batch_size 16
5. RLHF — DPO / GRPO
# DPO 训练(直接对比 chosen/rejected 样本,不需要奖励模型)
python scripts/dpo.py --config configs/dpo/base.json
# GRPO(DeepSeek 的 Group Relative Policy Optimization,当前主流方法)
python scripts/grpo.py --config configs/grpo/base.json
典型适用场景
- 学生/研究者:想深入理解 LLM 训练每个环节的原理,不满足于"调用 API"层面
- 教育者:用作 LLM 相关课程的实践教材,每个章节有代码、有输出
- AI 工程师:想从零搭建自定义 LLM pipeline,而不是依赖第三方微调框架
- 开源爱好者:想找一个透明、无黑箱的 LLM 训练参考实现
坑与注意
-
r50k_base tokenizer 限制:项目使用 OpenAI GPT-3 同款的
tiktokenr50k_base,不是现代 LLM 常用的tiktokencl100k_base 或 SentencePiece。如果你计划将训练好的权重与其他库(如transformers)互操作,需要注意 token IDs 不会对齐,需要额外做 embedding projection。 -
多阶段配置容易混淆:仓库有两套配置系统——
config/config.py(Legacy 预训练用)和config/post_training_config.py+configs/*.json(新路径),两者不要混用。 -
smoke 测试务必先跑:每次进入新阶段前,先跑
configs/smoke/下的极小配置确认 pipeline 正确,再上真实数据。 -
WandB 日志需要账号:如果不安
wandb,训练日志不会上传,但这不影响本地训练,只需不装[train]extra 即可。 -
数学推理的 answer tag:GRPO 阶段要求模型输出格式为
>>>...<<<包裹的答案(用于自动提取 reward),这个格式在 SFT 阶段也需要保持一致,否则 reward 为 0。 -
作者在寻找 PhD:README 明确说明作者在找教职,仓库更新可能不稳定,大版本改动前请 check commit 历史。
与同类对比
| 仓库 | 语言 | 依赖层 | RLHF | 特点 |
|---|---|---|---|---|
| 本仓库 | Python/PyTorch | 纯底层手写 | DPO/PPO/GRPO 全覆盖 | 教学导向,步骤最透明 |
| nanoGPT (karpathy) | Python/PyTorch | 底层 | 无 | 更简洁但只覆盖预训练 |
| llm.c (karpathy) | C/CUDA | 极底层 | 无 | 纯 C 实现预训练,速度优先 |
| HuggingFace TRL | Python/PyTorch | transformer 库 | DPO/PPO/SFT | 工业级但隐藏细节 |
本仓库的最大差异化价值:手写了所有 RLHF 算法(DPO/PPO/GRPO),不依赖 trl 等高层封装,特别适合想理解 RLHF 内部机制的学习者。
一句话推荐结论
想真正从零搞懂 LLM 训练(预训练 → SFT → RLHF)而不是停留在"调用 API"层面,这个仓库是当前最透明、最全面的开源教程之一,强烈推荐配合 README 原文和文档站一起学习。
来源:GitHub README (https://github.com/FareedKhan-dev/train-llm-from-scratch)、https://fareedkhan-dev.github.io/train-llm-from-scratch/