CS336:从零手搓LLM
只用基础框架的基础功能,如何从脚本开始手搓transformer,训练一个自己的大模型?
基于 TinyStories 数据集,从 BPE Tokenizer 训练到 Transformer 文本生成,完整实现一个现代语言模型的所有基础组件。
目录
项目概览
维度 | 详情 |
|---|---|
数据集 | TinyStories(英文儿童故事语料) |
模型参数 | ~15M |
词表大小 | 1000(BPE) |
上下文长度 | 256 tokens |
层数 | 4 层 Transformer |
隐藏维度 | d_model=512 |
注意力头数 | 16 |
FFN 维度 | d_ff=1344 |
训练框架 | PyTorch,纯手写组件(不调用 |
整个项目的文件结构如下:
assignment1-basics/
├── cs336_basics/ # 核心实现库
│ ├── module.py # 模型组件(Linear、Embedding、Attention、Transformer 等)
│ ├── functional_imp.py # 基础函数(softmax、SiLU、交叉熵损失、梯度裁剪)
│ ├── optimizer_imp.py # AdamW 优化器
│ ├── run_get_batch_imp.py # 批次数据加载
│ ├── run_train_bpe_imp.py # BPE 训练算法
│ ├── get_tokenizer_imp.py # Tokenizer 编码/解码
│ ├── model_imp.py # 手动实现 forward pass(基于权重字典)
│ ├── serialization_imp.py # 模型保存与加载
│ └── pretokenization_example.py # 预分词示例
├── combine/ # 流水线脚本
│ ├── config.py # 统一配置
│ ├── train_bpe.py # 训练 BPE tokenizer
│ ├── run_tokenize.py # 文本 → token ID 批量转换
│ ├── train_model.py # 模型训练主入口
│ ├── generate_text.py # 自回归文本生成(推理)
│ └── utils.py # tokenizer 加载工具
└── tests/ # 单元测试(共 8 个测试文件)第一部分:BPE Tokenizer
1.1 BPE 训练算法
Byte Pair Encoding 的核心思想是:从字节级开始,反复合并最高频的相邻 token 对,直到达到目标词表大小。
实现位于 run_train_bpe_imp.py,核心流程:
flowchart TD
A["原始文本字节序列"] --> B["预分词 pre-tokenization<br/>用正则把文本拆成 chunks"]
B --> C["每个 chunk 拆成单字节<br/>初始化 pair 列表"]
C --> D["统计所有相邻 pair 的频率<br/>找出最高频 pair"]
D --> E{"词表大小够了?"}
E -->|否| F["合并该 pair:生成新 token<br/>加入 vocab 和 merges"]
F --> G["更新所有 chunk 的 pair 列表"]
G --> D
E -->|是| H["输出 vocab + merges"]这里有一个关键设计——预分词。用正则表达式把文本拆成独立的 chunk(单词、标点等),每个 chunk 内部独立做 BPE 合并,跨 chunk 不合并。这避免了合并跨越语义边界。
1.2 训练细节
通过 train_bpe.py 调用训练:
vocab, merge = run_train_bpe_imp.run_train_bpe_imp(path, 1000, special_tokens)训练产物:
vocab.json:{token_id: bytes}的映射,词表大小 1000merge.json:合并规则列表,按优先级排序[(token_a, token_b), ...]
1.3 Tokenizer 编码与解码
实现位于 get_tokenizer_imp.py,编码过程:
文本 → UTF-8 字节 → 单字节 list → 贪婪合并(按 merges 优先级)→ token IDs
解码过程:token IDs → 查 vocab 得到 bytes → 拼接 → UTF-8 解码回文本
1.4 批量化处理
run_tokenize.py 将整个 TinyStories 数据集编码为 np.int16 二进制文件(int16 足够覆盖 1000 词表,比 int64 节省 4 倍空间)。
第二部分:神经网络基础组件
所有组件都在 module.py 中手写实现,不调用 nn.Linear 等高级 API。
2.1 Linear(线性层)
class Linear(torch.nn.Module):
def forward(self, x):
return einsum(x, self.weight, '... d, o d -> ... o')无偏置(
bias=None),简化设计权重形状
[out_features, in_features],与 PyTorch 原生一致使用
einsum实现,...自动适配任意 batch 维度权重初始化:
N(0, 0.02²)
2.2 Embedding(词嵌入)
class Embedding(torch.nn.Module):
def forward(self, token_ids):
return self.weight[token_ids]
# [batch, seq_len] 整数 ID → [batch, seq_len, d_model] 浮点向量本质上是一个可学习的查找表,用 PyTorch 的高级整数索引(weight[token_ids])实现。每个 token ID 被替换成对应行的高维向量。
2.3 RMSNorm(均方根归一化)
与 LayerNorm 的区别在于不做去均值(centering),计算量更小。LLaMA 系列模型就使用 RMSNorm。
class Rmsnorm(torch.nn.Module):
def forward(self, x):
x_fp32 = x.to(torch.float32) # 转高精度防溢出
rms = torch.sqrt(x.pow(2).sum(-1)/self.d_model + self.eps)
return (x_fp32 / rms * self.weight).to(orig_dtype)2.4 交叉熵损失(CrossEntropyLoss)
位于 functional_imp.py,使用了数值稳定的 log-sum-exp trick:
target_logits = torch.gather(inputs, -1, targets.unsqueeze(-1))
max_inputs = torch.max(inputs, dim=-1, keepdim=True).values
n_log = -target_logits + max_inputs + torch.log(torch.sum(exp(inputs - max_inputs)))
return n_log.mean()核心公式:
2.5 其他辅助函数
组件 | 作用 |
|---|---|
| 沿指定维度归一化为概率 |
| Swish 激活函数: |
| 按全局 L2 范数裁剪梯度,阈值 0.1 |
第三部分:Transformer 模型
3.1 模型整体结构
Transformer
├── token_embeddings: Embedding(1000, 512) # Token → 向量
├── layers: ModuleList[4 × TransformerBlock]
│ └── TransformerBlock
│ ├── ln1: RMSNorm # Pre-Norm ①
│ ├── attn: MultiheadAttentionWithRope # RoPE 多头注意力
│ ├── ln2: RMSNorm # Pre-Norm ②
│ └── ffn: SwiGLU(d_model=512, d_ff=1344) # SwiGLU 前馈网络
├── ln_final: RMSNorm # 最终归一化
└── lm_head: Linear(512, 1000) # 输出投影 → 词表分布架构选择:Pre-Norm + RoPE + SwiGLU,这是现代 GPT 系列(GPT-2/3、LLaMA)的标准配置。
3.2 Multi-Head Attention with RoPE
flowchart LR
subgraph 单个 TransformerBlock
A["输入 [B, T, 512]"] --> B["ln1 (RMSNorm)"]
B --> C["Q/K/V 投影 (Linear)"]
C --> D1["Q: [B, 16, T, 32]"]
C --> D2["K: [B, 16, T, 32]"]
C --> D3["V: [B, 16, T, 32]"]
D1 --> E1["RoPE 旋转编码"]
D2 --> E2["RoPE 旋转编码"]
E1 --> F["Q @ K^T / sqrt(32)"]
E2 --> F
F --> G["Causal Mask"]
G --> H["Softmax → Attention Weights"]
H --> I["Weighted Sum with V"]
D3 --> I
I --> J["Concat Heads (16×32→512)"]
J --> K["Output Projection"]
K --> L["+ Residual"]
end关键实现细节:
拆分多头:用 einops 将 [B, T, 512] 重排为 [B, 16, T, 32]
q_split = rearrange(q, "... s (h k) -> ... h s k", h=16)Casual Mask:动态根据实际序列长度创建下三角 mask
causal = torch.tril(torch.ones(seq_len, seq_len)).bool()
scores = scores.masked_fill(~causal, float("-inf"))合并多头:反转拆分,将 [B, 16, T, 32] 恢复为 [B, T, 512]
rearranged_head = rearrange(head, "... h s d -> ... s (h d)")3.3 RoPE(旋转位置编码)
RoPE 的核心是对特征维度做两两配对的 2D 旋转,不同 pair 使用不同频率:
key_pair = rearrange(x, "... (k p) -> ... k p", p=2) # 512维 → 256对×2
theta_p_i = pos × θ^(-2i/d) # 每对独立的旋转角度
x1, x2 = key_pair[..., 0], key_pair[..., 1]
y1 = x1 * cos - x2 * sin # 2D 旋转
y2 = x1 * sin + x2 * cos高频对 (i 小):旋转快,对近邻位置变化敏感
低频对 (i 大):旋转慢,适合编码长距离位置关系
这样 RoPE 用旋转矩阵的性质自然地编码了绝对和相对位置信息。
3.4 SwiGLU FFN
class Swiglu:
def forward(self, x):
return self.w2( self.w3(x) * silu(self.w1(x)) )公式:
相比传统 ReLU FFN(两个矩阵),SwiGLU 多了一个门控矩阵 W3,用 SiLU 激活做门控,表达能力更强。
3.5 前向传播流程
[token IDs] → Embedding → [Pre-Norm → Attention → +Residual
→ Pre-Norm → SwiGLU → +Residual] × 4
→ Final RMSNorm → Linear → [logits]我们还在模型层增加了输入长度自动截断保护:
def forward(self, in_indices):
if in_indices.shape[-1] > self.context_length:
in_indices = in_indices[..., -self.context_length:] # 保留最后 256 token
...这样调用方无需关心上下文长度限制,模型内部自动处理。
第四部分:训练与推理
4.1 优化器:AdamW
从头实现 AdamW(位于 optimizer_imp.py):
m_t = β₁·m_{t-1} + (1-β₁)·grad # 一阶动量
v_t = β₂·v_{t-1} + (1-β₂)·grad² # 二阶动量
lr_t = lr · √(1-β₂^t) / (1-β₁^t) # 偏差修正
θ -= lr_t · m_t / (√v_t + ε) # Adam 更新
θ -= lr · weight_decay · θ # Weight Decay(解耦,AdamW 关键改进)超参数:lr=3e-4, β=(0.9, 0.95), weight_decay=0.01, ε=1e-8
4.2 数据加载
使用 np.memmap 以内存映射方式读取二进制的 token ID 文件——不占用内存,按需从磁盘读取:
train_token_ids = np.memmap("train.bin", dtype=np.int16, mode="r")每个 batch 的训练步骤:
# 随机采样 32 个不重复的起始位置
begin_indexes = random.sample(range(0, n - 256), 32)
x = [dataset[i : i+256] for i in begin_indexes] # 输入序列
y = [dataset[i+1 : i+257] for i in begin_indexes] # 标签(右移一位)这是标准的自回归语言建模方式:输入 tokens[0:256],预测 tokens[1:257]。
4.3 训练循环
flowchart TD
subgraph 每个 batch
A["随机采样 (x, y)"] --> B["zero_grad"]
B --> C["model.forward → logits [32,256,1000]"]
C --> D["CrossEntropyLoss → loss"]
D --> E["loss.backward"]
E --> F["梯度裁剪 (max_l2_norm=0.1)"]
F --> G["optimizer.step"]
end
G --> H{iter % 2000 == 0?}
H -->|是| I["验证集评估"]
H -->|是| J["保存 checkpoint"]每 2000 步做一次验证集评估并保存 checkpoint,每 100 步打印训练 loss。
4.4 文本生成(推理)
generate_text.py 实现了自回归生成:
flowchart TD
A["给定 prompt 文本"] --> B["tokenizer.encode → token IDs"]
B --> C["model.forward → logits [T, 1000]"]
C --> D["取最后一个位置: logits[-1, :]"]
D --> E["temperature 缩放"]
E --> F{使用 top_k?}
F -->|top_k=50| G["取 top-50 概率最高 token"]
F -->|否| H["全词表"]
G --> I["softmax → 概率分布"]
H --> I
I --> J["torch.multinomial 随机采样"]
J --> K["tokenizer.decode → 文本"]
K --> L{"遇到结束符?"}
L -->|否| B
L -->|是| M["结束"]采样策略:top-k = 50 + temperature。从概率最高的 50 个 token 中按概率随机选一个,避免低概率垃圾 token,同时保持生成多样性。
4.5 Bug 修复记录
在开发过程中修复了一些问题:
问题 | 修复 |
|---|---|
推理时按字符截断而非按 token 截断 | 改为先编码再按 token 数截断 |
截断逻辑放在调用侧,职责不清 | 移到 |
| 配置 |
踩坑记录与经验教训
1. np.memmap 是处理大数据集的好帮手
TinyStories 训练集 token 化后大约几百 MB,传统方式加载会撑爆内存。np.memmap 映射到虚拟地址空间,按需读取,完美解决。
2. einsum 让张量操作一目了然
# 传统 torch
output = torch.matmul(x, weight.T)
# einops
output = einsum(x, weight, '... d, o d -> ... o')后者读起来更直观,... 自动适配任意 batch 维度,省去大量 reshape 操作。
3. 数值稳定性很重要
CrossEntropyLoss 的 log-sum-exp trick、RMSNorm 转 float32 计算——这些都是保证训练不会 NaN 的关键细节。
4. RoPE 不检查序列长度
模型中 max_seq_len 参数被存储但从未在 forward 中校验。超过训练时见过的位置范围(256)不会报错,但 RoPE 编码的位置模型从未学过,会导致外推质量严重下降。
5. 职责归属要清晰
推理时应该在模型层做截断而不是在调用侧,这样所有调用方都能自动受保护,不需要各自关心内部限制。
总结
通过这次作业,从 BPE tokenizer 训练到 Transformer 模型实现的完整链路都走过了一遍。核心收获:
BPE 算法:理解了子词切分的原理和实现细节
现代 Transformer 架构:Pre-Norm、RoPE、SwiGLU 这些设计选择背后的动机
从零实现:Linear、Embedding、RMSNorm、Attention、AdamW 等组件的手写实现加深了对底层计算的理解
工程实践:
np.memmap处理大数据、einops简化张量操作、数值稳定性的重要性调试经验:字符 vs token 截断、模型层 vs 调用侧的职责划分
最终模型基于 TinyStories 数据集训练的这约 1500 万参数的小型 Transformer,能够生成类似儿童故事的连贯文本。
本文基于 CS336 (Stanford) 2025 春季学期 Assignment 1 的实战记录。完整代码见项目仓库。

Comments
Discuss this project
Emoji supported. Comments appear immediately.