如何手动构建一个LLM

如何手动构建一个 LLM

最近在看karpathy/nanochat 的源码,试着跟着项目把"从零搭一个 LLM"这件事做一遍。这里记录一下整体流程。

nanochat是什么

nanochat是karpathy基于nanoGPT这个项目的全新改进版本,覆盖截止2025年底大模型全生命周期内训你的步骤的一个极简版本。

Karpathy 的思路很直接:用尽量少的代码,实现一个LLM。整个项目核心库只有十几个 .py 文件,包含了训练一个LLM所必须的所有步骤:

  • 分词(Tokenizer)训练
  • 基座(Base)自回归预训练
  • 监督微调(SFT)
  • 强化学习(RL / GRPO 风格)
  • 交互式推理与 Tool Calling

完成这些步骤,在这个项目的底层只靠一个复杂度参数进行驱动:--depth(Transformer 层数)。通过设置这个参数,后续步骤中需要用到的:模型宽度、学习率、批大小等参数就按缩放律自动算出来,不用再去翻配置表纠结。

ps:缩放律指的是模型参数量、训练数据量、计算量,与模型能力(loss、准确率等)之间的关系,目前这个关系是数据量越大,模型能力越强

二、技术栈

层次 选型
语言 / 包管理 Python 3 + uv
深度学习框架 PyTorch 2.9.1
分布式训练 手动 DDP(不依赖 torch.nn.parallel.DistributedDataParallel),ZeRO-2 风格梯度分片
注意力 Flash Attention 3(FA3)+ SDPA 回退
量化训练 自研 FP8(约 150 行,替代 torchao
分词器 RustBPETokenizer(训练用 rustbpe,推理用 tiktoken
优化器 自研 MuonAdamW(2D 参数走 Muon,1D 参数走 AdamW)
数据 PyArrow Parquet + BOS-aligned best-fit packing
监控 Weights & Biases(可选)
编译 torch.compiledynamic=False
推理引擎 自研 KVCache + Python Tool Calling 状态机

三、训练的五个步骤

项目的入口脚本串成一条主线,主要包含一下五个步骤:

tok_train → base_train → chat_sft → chat_rl → chat_cli
    ↓           ↓            ↓           ↓          ↓
 BPE 分词器   预训练基座    监督微调     强化学习    交互对话
  1. tok_train — 使用统计方法训练自己的 BPE 分词器,生成词表
  2. base_train — GPT 自回归语言模型预训练(核心流程)
  3. chat_sft — 在对话数据(SmolTalk + MMLU + GSM8K)上做监督微调
  4. chat_rl — GRPO 风格 RL(简化成 REINFORCE),优化数学推理
  5. chat_cli — 交互式对话 CLI,支持 Tool Calling

四、Stage 1 — 先训一个分词器

大模型训练的第一个步骤,就是把文本切分成token。nanochat 用 RustBPETokenizer

  • 训练阶段用 rustbpe 把原始语料按照合并规则,用 9 个特殊 token进行分割,这9个特殊的Token分别是1(文档边界)+ 4(对话角色)+ 4(工具调用)的token,剩下的 token 就是 BPE 词表

在第二个步骤的推理阶段会使用 tiktoken 做高效的编解码,将文本↔id

这个步骤是后面所有步骤的基础。模型所看到的每个token,都是它分割出来的,也就是说后续所有的步骤都是使用这个阶段切分出来的语料。

rustbpe 是一个用 Rust 写的 BPE 分词器训练库(PyPI 上的 rustbpe 包)。在 nanochat 里,它只干一件事:从零训练分词器
tiktoken 是 OpenAI 开源的高性能分词库(GPT-4 用的就是它)。在 nanochat 里,它负责推理阶段的编解码。

五、Stage 2 — 基座预训练(base_train.py)

这个是LLM训练的核心,目标很简单:用一份语料,训练出一个能预测下一个 token 的模型。这个阶段的训练是自回归的,模型的输入是前面所有的 token,输出是下一个 token 的预测。

1. Meta Device 初始化(躲开 CPU 显存峰值)

model = GPT(config)          # 在 meta 设备上画蓝图,不占显存
model = model.to_empty(device="cuda")  # 真正分配显存
model.init_weights()         # 按缩放律初始化

这个步骤做的事情是
第一:使用PyTorch的Meta Device功能根据模型的配置,构造出神经网络训练所需的结构和参数;
第二:使用to_empty方法将模型的参数分配到GPU上,避免在初始化时占用过多显存;
第三:使用init_weights方法根据缩放律初始化模型的权重参数;

2. 自定义 Linear 层(手动混合精度)
不靠 torch.amp.autocast,而是自己控制:

class Linear(nn.Module):
    def forward(self, x):
        return F.linear(x, self.weight.to(x.dtype))  # fp32 主权重,按输入 dtype(bf16) 算

权重留 fp32(master weights),前向时 cast 到 bf16 计算,这就是手动混合精度的核心。

这个步骤主要是做两件事情:
第一件事情是进行矩阵乘法,这一步是计算神经网络中最重要的操作之一,它将输入数据与权重矩阵相乘,得到输出结果,这里的输出结果指的是神经网络的预测结果,在这个过程中,权重矩阵是以 fp32 的精度存储的,而输入数据通常是以 bf16 的精度存储的,为了保证计算的精度和效率,在进行矩阵乘法之前,需要将权重矩阵从 fp32 转换为输入数据的 dtype(通常是 bf16),这样可以在保证计算精度的同时,提高计算效率和减少显存占用;
第二件事情是进行数据类型转换,在计算过程中,将权重矩阵从 fp32 转换为输入数据的 dtype(通常是 bf16), 这样可以在保证计算精度的同时,提高计算效率和减少显存占用

用一个简单的例子来说就是我们输入一段话“小明喜欢吃”,然后模型先把这段话转化为向量(Stage 1),Stage 2的这个步骤就是权重矩阵相乘,得到一个新的向量,这个新的向量就是模型的预测结果。为了保证计算的精度和效率,我们需要把权重矩阵从 fp32 转换为输入数据的 dtype(通常是 bf16),这样可以在保证计算精度的同时,提高计算效率和减少显存占用。
这里使用bf16是因为它在计算速度和显存占用上有优势,而且在大多数情况下,它的精度已经足够满足模型的需求。

3. MuonAdamW 混合优化器

  • 2D 参数(权重矩阵)→ Muon:Polar Express 正交化(5 次 Newton-Schulz 迭代)+ MuonEq + Muon+ + NorMuon,加速收敛
  • 1D 参数(embedding / bias / scalar)→ AdamW(fused kernel)

梯度同步走三阶段异步,替代 PyTorch DDP:

Phase 1: reduce_scatter 梯度  →  AllReduce
Phase 2: 本地更新参数          →  compute
Phase 3: all_gather 同步权重   →  finish

计算和通信重叠,把带宽吃满。

4. FP8 量化训练
自研 FP8 也就一百来行,替代 torchao:前向用 e4m3(fast accum),反向用 e5m2,靠 torch._scaled_mm 落地。

5. BOS-aligned Best-Fit Packing 数据加载
按 BOS 对齐切分文档,用 best-fit 装箱把多篇塞进定长序列(利用率接近 100%,裁剪率约 35%);每个 rank 读自己的 Parquet shard,预分配 pinned CPU/GPU buffer,单次 HtoD 传输。

6. 缩放律自动推导
tokens = 12 × paramsB ∝ D^0.383、学习率 η ∝ √(B/B_ref)、weight decay λ ∝ √(B/B_ref)·(D_ref/D),全由 depth 锚定。

六、Stage 3 — 监督微调(chat_sft.py)

加载基座 checkpoint,在它那套超参上做对话微调:

  • TaskMixture:SmolTalk + MMLU×epochs + GSM8K×epochs 混合采样
  • 数据生成用 padding 而非裁剪,mask=-1 忽略 padding 位置
  • LR 调度看进度比例(0→1),而不是绝对步数

产出的是第一个"能聊天"的模型。

七、Stage 4 — 强化学习(chat_rl.py)

GRPO 的极简版:

  • advantages = rewards - mu(不做 z-score 归一化)
  • pg_obj = (logp * advantages).sum() / num_valid
  • LR 线性 rampdown 到零
  • pass@k 评估推理可靠性

用 GSM8K 数学 rollout 来练模型的逻辑推理。

八、Stage 5 — 推理与对话(engine.py + chat_cli.py)

  • KVCache:FA3 原生布局 (B,T,H,D)prefill → clone KV → decode loop 模式
  • Tool Calling 状态机RowState):
    user → assistant_gen → [python_start → code → python_end]
    → [output_start → result → output_end] → assistant_cont → eos
    <|python_start|>…<|python_end|> 触发沙箱代码执行,结果经 <|output_start|>…<|python_end|> 注回上下文
  • 沙箱执行execution.py):subprocess 隔离 + 禁掉危险 builtins + 内存 rlimit(256MB)+ tempdir 隔离

九、核心模型架构(gpt.py 详解)

这是项目的心脏(约 800 行),几个关键设计:

  • RoPE 位置编码
  • QK Norm(×1.2):注意力的稳定器
  • ValueEmbed:ResFormer 风格,value 投影前先过 embedding gate
  • ReLU² MLP:替代 GELU/SwiGLU 的极简激活
  • Logit Softcap = 15.0:限制 logit 幅值,稳住训练
  • 滑动窗口注意力(SSSL pattern)
  • Smear:embedding 初始化时把 logits 分布均匀涂抹
  • Backout:训练不稳时回退权重
  • padded_vocab_size = ((vocab_size + 63) // 64) * 64(对齐到 64)

前向流:

input_ids → tok_embed + ValueEmbeds → Smear
→ [Block×depth: RMSNorm → Attn(RoPE+QKNorm+ValueEmbed+滑窗) → RMSNorm → MLP(ReLU²)]
→ RMSNorm → lm_head(untied) → logits → Softcap → cross_entropy

十、工程取舍与设计模式

模式 体现 评价
单一复杂度旋钮 --depth 驱动一切 优雅,降低门槛
策略模式 Muon/AdamW 按维度分流 各取所长
模板方法 Task / TaskMixture / TaskSequence 数据组合清晰
状态机 RowState(推理 Tool Calling) 清晰可控
Builder 模式 build_model_meta() 参数化构建
手动 DDP 替代 PyTorch DDP 包装 控制更细,但更复杂
GC 管理策略 首步后 disable + 定时 collect 避免 GC 抖动

十一、技术债务与优化路线

技术债务:

  1. 安全沙箱不充分(execution.pyexec() 禁用可被熟手绕过)
  2. 手动 DDP 维护成本高(三阶段异步通信逻辑复杂,得跟着 PyTorch 版本走)
  3. 硬编码路径与常量(魔数散落,缺统一 config 入口)
  4. 只支持单机多卡(手动 DDP 不跨节点)
  5. FP8 实现偏简(缺动态缩放、NaN 检测、梯度裁剪)
  6. 测试覆盖薄(仅 6 个测试文件,缺优化器/分布式/精度回归)
  7. Tokenizer 双轨依赖(训练 rustbpe / 推理 tiktoken,得保证行为一致)

优化建议:

  • 高优先级:把 execution.py 换成 Docker/gVisor 级沙箱;给 FP8 补上数值稳定性保障
  • 中优先级:手动 DDP 迁到 FSDP2 + torch.compile;超参抽到统一 YAML;补 CI 测试
  • 低优先级:支持多节点训练;Tokenizer 统一到单一后端;加 ANNX/GGUF 导出

结语

跑完 chat_cli 之后,看着模型顺着提示流畅续写、调用工具、一步步推数学题,我算是彻底信了:一个 LLM 真没那么神秘。它不过是一堆精心设计的 Linear、Attention、RoPE,加上一次三阶段梯度同步的组合。

nanochat 把这份"神秘感"拆成了能逐行读、逐阶段改的约 2000 行代码。亲手搭起来之后,它不再是个黑盒。

完整源码与运行说明见 karpathy/nanochat

参考资料

Deep Dive into LLMs like ChatGPT