三藏签名
< Back to projects

CS336:从零手搓LLM

AILLM

只用基础框架的基础功能,如何从脚本开始手搓transformer,训练一个自己的大模型?

基于 TinyStories 数据集,从 BPE Tokenizer 训练到 Transformer 文本生成,完整实现一个现代语言模型的所有基础组件。


目录

  1. 项目概览

  2. 第一部分:BPE Tokenizer

  3. 第二部分:神经网络基础组件

  4. 第三部分:Transformer 模型

  5. 第四部分:训练与推理

  6. 踩坑记录与经验教训

  7. 总结


项目概览

维度

详情

数据集

TinyStories(英文儿童故事语料)

模型参数

~15M

词表大小

1000(BPE)

上下文长度

256 tokens

层数

4 层 Transformer

隐藏维度

d_model=512

注意力头数

16

FFN 维度

d_ff=1344

训练框架

PyTorch,纯手写组件(不调用 nn.Linear 等高级 API)

整个项目的文件结构如下:

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} 的映射,词表大小 1000

  • merge.json:合并规则列表,按优先级排序 [(token_a, token_b), ...]

1.3 Tokenizer 编码与解码

实现位于 get_tokenizer_imp.py,编码过程:

  1. 文本 → UTF-8 字节 → 单字节 list → 贪婪合并(按 merges 优先级)→ token IDs

  2. 解码过程: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(均方根归一化)

RMSNorm(x)=x1dxi2+ϵγ\text{RMSNorm}(x) = \frac{x}{\sqrt{\frac{1}{d}\sum x_i^2 + \epsilon}} \cdot \gamma

与 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()

核心公式:

L=1Ni[zi,yimax(zi)logjezi,jmax(zi)]\mathcal{L} = -\frac{1}{N}\sum_i \left[ z_{i,y_i} - \max(z_i) - \log\sum_j e^{z_{i,j} - \max(z_i)} \right]

2.5 其他辅助函数

组件

作用

softmax

沿指定维度归一化为概率

SiLU

Swish 激活函数:x · σ(x)

gradient_clipping

按全局 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)) )

公式:FFN(x)=W2(SiLU(W1x)W3x)\text{FFN}(x) = W_2(\text{SiLU}(W_1x) \odot W_3x)

相比传统 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 数截断

截断逻辑放在调用侧,职责不清

移到 Transformer.forward() 内部自动处理

config 导入 Pylance 报错

配置 .vscode/settings.json 添加 extraPaths


踩坑记录与经验教训

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. 数值稳定性很重要

CrossEntropyLosslog-sum-exp trick、RMSNormfloat32 计算——这些都是保证训练不会 NaN 的关键细节。

4. RoPE 不检查序列长度

模型中 max_seq_len 参数被存储但从未在 forward 中校验。超过训练时见过的位置范围(256)不会报错,但 RoPE 编码的位置模型从未学过,会导致外推质量严重下降。

5. 职责归属要清晰

推理时应该在模型层做截断而不是在调用侧,这样所有调用方都能自动受保护,不需要各自关心内部限制。


总结

通过这次作业,从 BPE tokenizer 训练到 Transformer 模型实现的完整链路都走过了一遍。核心收获:

  1. BPE 算法:理解了子词切分的原理和实现细节

  2. 现代 Transformer 架构:Pre-Norm、RoPE、SwiGLU 这些设计选择背后的动机

  3. 从零实现:Linear、Embedding、RMSNorm、Attention、AdamW 等组件的手写实现加深了对底层计算的理解

  4. 工程实践np.memmap 处理大数据、einops 简化张量操作、数值稳定性的重要性

  5. 调试经验:字符 vs token 截断、模型层 vs 调用侧的职责划分

最终模型基于 TinyStories 数据集训练的这约 1500 万参数的小型 Transformer,能够生成类似儿童故事的连贯文本。


本文基于 CS336 (Stanford) 2025 春季学期 Assignment 1 的实战记录。完整代码见项目仓库。

Comments

Discuss this project

Emoji supported. Comments appear immediately.

No comments yet.