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 实现,不依赖 transformerstrlpeft 等高层库,从文本语料处理讲起,一路走到 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 训练参考实现

坑与注意

  1. r50k_base tokenizer 限制:项目使用 OpenAI GPT-3 同款的 tiktoken r50k_base,不是现代 LLM 常用的 tiktoken cl100k_base 或 SentencePiece。如果你计划将训练好的权重与其他库(如 transformers)互操作,需要注意 token IDs 不会对齐,需要额外做 embedding projection。

  2. 多阶段配置容易混淆:仓库有两套配置系统——config/config.py(Legacy 预训练用)和 config/post_training_config.py + configs/*.json(新路径),两者不要混用。

  3. smoke 测试务必先跑:每次进入新阶段前,先跑 configs/smoke/ 下的极小配置确认 pipeline 正确,再上真实数据。

  4. WandB 日志需要账号:如果不安 wandb,训练日志不会上传,但这不影响本地训练,只需不装 [train] extra 即可。

  5. 数学推理的 answer tag:GRPO 阶段要求模型输出格式为 >>>...<<< 包裹的答案(用于自动提取 reward),这个格式在 SFT 阶段也需要保持一致,否则 reward 为 0。

  6. 作者在寻找 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/