CS336 Assignment 1:从零实现一个小型语言模型

Author

jshn9515

Published

2026-07-17

Modified

2026-07-17

这一部分对应 CS336 Assignment 1:从零实现一个小型语言模型。

Assignment 1 的核心目标是把前面 Chapter 8 和 Chapter 19 中介绍的 Transformer、attention、BPE Tokenizer 等模块真正组合起来,实现一个完整的 GPT 风格语言模型训练流程。相比前面的理论介绍,这部分更偏工程实践,包括 tokenizer、Transformer block、训练循环、优化与评估等完整流程。

本部分包含 Assignment 1 总结,对作业中的各个问题进行了实现说明、实验记录和结果分析。

所有核心组件均实现于 dnnlpy 中,包括:

所有模型实现独立于 torch.nn API,复用前面笔记中已经实现的深度学习基础组件。

通过这一部分,可以进一步理解现代 LLM 的基本训练过程:

后续章节会在此基础上继续介绍 LLM 训练工程、数据处理、scaling law、推理优化以及后训练方法。

from collections import defaultdict

import dnnlpy
import dnnlpy.models.gpt as gpt
import dnnlpy.nn as dnn
import dnnlpy.nn.functional as dF
import dnnlpy.optim as dopt
import dnnlpy.tokenizers as dltk
import IPython.display as ipy
import pandas as pd
import torch
import torch.nn as nn

print('PyTorch version:', torch.__version__)
PyTorch version: 2.13.0+xpu

2.1 Unicode

Problem (unicode1)

  1. chr(0) 返回的是空字符(null character),对应的 Unicode 编码点为 U+0000.
  2. 它的 repr 表示形式是可见的转义序列 \x00;但是直接使用 print() 输出时,会发送一个不可见的控制字符,因此通常不会在终端中显示任何内容。
  3. 该字符可以存在于 Python 字符串内部,并且会被计入字符串长度。虽然它在显示时通常不可见,但它仍然是字符串中的一个有效字符。
null = chr(0)
print('repr(null):', repr(null))
print('len(null):', len(null))
print('ord(null):', ord(null))
repr(null): '\x00'
len(null): 1
ord(null): 0

Problem (unicode2)

  1. UTF-8 通常是更优的选择,因为它是基于字节的编码方式,对于以 ASCII 字符为主的文本更加紧凑,在 Web 环境中占据主导地位,并且不会像 UTF-16/UTF-32 在英文文本中那样频繁产生大量的零字节。
  2. 如果逐个字节独立解码,会导致多字节字符解码失败。例如:'é'.encode() = b'\xc3\xa9'。但是,单独的 b'\xc3' 是一个不完整的 UTF-8 序列,无法被正确解码。
  3. b'\xc3\x28' 是一个无效的 UTF-8 编码序列:0xc3 表示一个双字节 UTF-8 字符的起始字节,但是后面的 0x28 并不是合法的 UTF-8 续接字节(continuation byte)。因此,该字节序列无法被解析为有效的 UTF-8 字符。
encoded = 'é'.encode('utf-8')

try:
    [bytes([byte]).decode('utf-8') for byte in encoded]
except UnicodeDecodeError as err:
    print('UnicodeDecodeError:', err)

try:
    b'\xc3\x28'.decode('utf-8')
except UnicodeDecodeError as err:
     print('UnicodeDecodeError:', err)
UnicodeDecodeError: 'utf-8' codec can't decode byte 0xc3 in position 0: unexpected end of data
UnicodeDecodeError: 'utf-8' codec can't decode byte 0xc3 in position 0: invalid continuation byte

2.4 BPE Tokenizer Training

Problem (train_bpe)

dnnlpy 状态:已在 dnnlpy.tokenizers.BPETrainer 中实现。

首先,实现一个最简单的 BPE。这个版本只关注算法流程,不考虑任何优化。在每一次 merge 操作中,直接使用 max() 函数遍历所有候选 pair,选择当前频率最高的 pair。同时,这个实现没有维护 pair 的倒查表,也没有使用多进程、多线程或其他缓存机制。因此,在合并一个 pair 后,需要重新遍历整个词表,找到所有包含该 pair 的 word,并更新对应的统计信息。这种实现虽然简单直观,但大量重复扫描会带来较高的计算开销。

假设进行 \(M\) 次 merge,每次需要扫描 \(P\) 个候选 pair,并遍历 \(V\) 个 word 更新统计,每个 word 的平均长度为 \(L\),那么整体计算的复杂度约为:

\[ O(M \cdot (P + VL)) \]

这在大规模语料和较大词表的情况下,训练时间是不可接受的。因此,需要引入一系列优化策略来提升训练效率。

为了确定这个简单实现的性能瓶颈,首先使用 Python 自带的 cProfile 工具进行 profiling。相比直接分析代码哪里慢,profiling 可以统计每个函数的调用次数和累计运行时间,从而定位真正影响性能的部分。

Tip

Python 3.15 将引入新的 profiling 模块,内部使用 Tachyon profiler,提供更高精度的性能分析和更低的性能开销。详细内容请参考 Python 3.15 profiling 模块文档

例如,可以通过:

python -m cProfile -s cumulative train_bpe.py

对 BPE 训练过程进行分析。其中:

  • ncalls 表示函数被调用的次数;
  • tottime 表示函数自身的运行时间,不包括调用子函数的时间;
  • cumtime 表示函数及其子函数的累计运行时间。

运行 profile 后,可以发现大量时间消耗在以下几个步骤:

  1. 寻找最高频 pair。每一次 merge 都需要调用 max() 遍历所有候选 pair,当 pair 数量较大时,这个过程会被重复执行很多次,导致性能瓶颈。
  2. 更新 pair 统计。由于没有维护倒查表,每次 merge 后都需要重新扫描整个词表,寻找包含目标 pair 的 word,并重新计算相关频率。这个步骤通常会占据大量运行时间,因为它涉及大量重复的数据遍历。

因此,总体优化思路是,首先优化 pair 查找过程,不再每次使用 max() 遍历所有 pair,而是维护一个能够快速获取最高频 pair 的数据结构。随后,再进一步优化 merge 后的更新过程,通过保存 pair 到 word 的倒查关系,只更新受到影响的部分。

第一轮优化:使用 heap 加速最高频 pair 查找

通过 profiling 可以发现,简单实现中一个明显的瓶颈是:每一次 merge 都需要调用 max() 遍历所有候选 pair,寻找当前频率最高的 pair。随着 pair 数量增加,这一步会被重复执行大量次数。实际上,在整个 BPE 训练过程中,很多时间都花在了反复寻找最大值上,而不是 merge 本身。

因此,第一步优化是引入 heap(优先队列)。相比每次从全部 pair 中搜索最大值,heap 可以提前维护当前频率最高的 pair,使得获取最高频 pair 的操作更加高效。具体来说,将 pair 按照 frequency 分桶,并维护一个 max-heap。每次需要选择下一个 merge pair 时,直接从 heap 顶部获取当前频率最高的 pair,而不需要重新扫描所有候选 pair。

但是,在 merge 过程中,pair 的 frequency 会不断变化。例如,合并一个 pair 后,相关 word 中的其他 pair 频率可能增加或减少。如果每次更新时都立即在 heap 中寻找并删除旧的记录,删除操作本身仍然需要遍历 heap,反而引入新的开销(我们知道,在 heap 中删除任意元素的最坏时间复杂度是 \(O(N)\),因为需要先找到该元素的位置)。

因此,这里采用延迟删除(lazy deletion)。也就是说,当一个 pair 的频率变化时,不立即删除 heap 中的旧记录,而是直接加入新的记录。这样 heap 中可能暂时存着一些已经过期的数据。当之后从 heap 顶部取出一个 pair 时,再检查它的 frequency 是否仍然是最新的。如果不是,就说明这个记录已经过期,直接丢弃,继续查看下一个 pair。这样,我们可以把删除操作的开销分摊到后续的 pop 操作中,时间复杂度仍然保持在 \(O(\log N)\)

Python 3.14 已经提供了原生的 max-heap API,例如 heapify_maxheappush_maxheappop_max。在旧版本 Python 中,可以通过保存负 frequency 的方式,使用普通 min-heap 模拟 max-heap。

经过这一轮优化后,pair 查找从每次 merge 都扫描全部 pair,变成了通过 heap 快速获取候选 pair。接下来继续使用 profiling,可以观察新的瓶颈,并进一步优化 pair 更新过程。

第二轮优化:只更新受 merge 影响的 word

第一轮优化解决了如何快速找到最高频 pair 的问题。但是 profiling 后可以发现,训练过程中的另一个主要瓶颈仍然存在:每次 merge 后,都需要重新扫描大量 word 来更新 pair frequency。实际上,一次 merge 只会影响包含当前 pair 的 word。对于不包含这个 pair 的 word,它们的 token 序列完全没有变化,对应的 pair frequency 也不会发生改变。因此,没有必要在每轮 merge 后重新遍历整个语料。

为了解决这个问题,在初始化阶段先完整统计一次 pair 信息,并额外维护 pair 到 word 的反向索引:

  • pair_counts[pair] -> frequency:保存每个 pair 当前的总频率;
  • pair_indices[pair] -> set[word_index]:记录哪些 word 包含该 pair。

当选择出当前需要 merge 的 best_pair 后,可以通过 pair_indices[best_pair] 直接找到所有受影响的 word,只对这些 word 执行 merge 和统计更新。

对于每个受影响的 word,只需要比较 merge 前后的局部变化,并更新发生改变的 pair frequency。假设某个 pair 在一个 word 中出现次数变化为:

\[ \Delta n(p) = \eta_{\mathrm{new}}​(p) - \eta_{\mathrm{old}}(p) \]

那么它对整体 frequency 的贡献变化为:

\[ \Delta c(p) = \Delta n(p) \times \mathrm{freq}(\mathrm{word}) \]

只有发生变化的 pair 才需要更新对应的统计信息。

通过这种方式,merge loop 从原来的“每轮扫描整个语料,重新统计所有 pair”变成了“只访问包含当前 merge pair 的 word,并增量更新受到影响的 pair”。虽然需要额外的内存来保存 pair 到 word 的倒查索引,但显著减少了重复扫描的开销。

第三轮优化:优化训练阶段的预分词过程

前两轮优化主要减少了 BPE merge loop 中存在的重复搜索和重复统计,但是 profiling 后仍然可以发现,训练开始之前的 pre-tokenization 阶段存在较大的额外开销。

最初的实现直接调用普通的 pre_tokenize() 流程。这个流程主要是为实际编码服务的,因此会生成完整的 (token, offset) 信息,并在每次遇到 token 时执行对应的 Unicode 转换。但是,对于 BPE 训练来说,我们只需要知道每个 token 出现的次数,并不需要保存 offset 信息。同时,在大规模语料中,相同的 token 往往会重复出现,如果每次都重新执行 Unicode 转换,会产生大量重复计算。

因此,我们为训练阶段增加一个专门的 fast path。这个路径不再生成 (token, offset) 对,而是直接统计原始 piece 的出现频率。随后,只需要对不同的 piece 执行一次 Unicode 转换,再将结果用于后续 BPE 训练。这里的 piece 指的是经过预分词(pre-tokenization)后得到的原始文本片段,还没有经过 BPE merge 的 token。

通过这一轮优化,训练阶段避免了两类不必要的工作:

  • 不再计算训练过程中不会使用的 offset 信息;
  • 不再对重复出现的 token 多次执行相同的 Unicode 转换。

相比前两轮针对 merge loop 的优化,这一次优化减少的是数据预处理阶段的额外开销,使整个 BPE 训练流程更加高效。

第四轮优化:批量并行预分词

经过前几轮优化后,merge loop 中的大量重复计算已经减少,但是 profiling 仍然显示,预分词(pre-tokenization)仍然占据了较多运行时间。

考虑到预分词过程中的每个 document 之间相互独立,因此天然适合并行处理。最直接的想法是为每个 document 创建一个独立的任务,交给多个 worker 执行。但是,实际测试发现,如果任务粒度过小,单独提交大量 document 任务会引入额外开销。例如,任务调度、对象序列化以及进程间数据传输都会消耗时间。当单个任务本身很短时,这些开销甚至可能超过并行计算带来的收益。

因此,我们将多个 document 合并成一个 batch,再把 batch 作为并行任务提交。这样可以减少任务数量,提高每个 worker 的工作量,同时降低调度和通信开销。具体来说,预分词阶段按照固定大小划分 batch:

  • 每个 batch 内包含多个 document;
  • 每个 worker 负责处理整个 batch;
  • 不同 batch 之间可以并行统计 token frequency;
  • 最后将各个 worker 返回的统计结果合并。

在实现上,根据 Python 运行环境选择不同的并行方式:

  • 在普通 Python 环境中,由于 GIL 限制,因此使用 ProcessPoolExecutor 创建多个进程执行;
  • 在支持自由线程的 Python 环境中,直接使用 ThreadPoolExecutor,避免进程创建和数据传输的额外开销;
  • num_workers 等于 1 时,直接使用串行路径,避免不必要的并行管理成本。

需要注意的是,并不是所有 BPE 训练过程都可以并行化。预分词只是统计每个 token 的出现频率,各 batch 之间相互独立,因此可以安全并行。而 BPE merge loop 中,每一次 merge 都依赖前一次 merge 更新后的 pair frequency,因此存在严格的数据依赖关系,仍然需要保持串行执行。

通过这一轮优化,主要减少了训练前数据处理阶段的时间开销,使整体训练流程更加接近实际 tokenizer 工程中的实现方式。

BPE 训练调用栈

从作业脚本开始,完整的训练调用栈如下。所有数据采用流式处理,避免一次性加载整个语料到内存中。

train_bpe_tinystories()
├─ datasets.load_dataset(...)
├─ _batch_iterator(dataset)
├─ Tokenizer.train_from_iterator(...)
│  └─ BPE.train_from_iterator(tokenizer, ...)
│     ├─ BPETrainer(...)
│     └─ BPETrainer.train(texts)
│        ├─ _iter_texts(texts)
│        ├─ _count_pre_tokens(texts)
│        │  ├─ Tokenizer._normalize(text)
│        │  └─ Tokenizer._count_pre_tokens(...)
│        │     ├─ it.batched(..., 1024)
│        │     ├─ parallel_map(...)
│        │     └─ ByteLevelPreTokenizer.count_pre_tokens(...)
│        ├─ _init_vocab(word_counts)
│        ├─ _init_pair_counts(word_symbols, word_freqs)
│        ├─ _init_pair_freq_index(pair_counts)
│        ├─ while len(vocab) < vocab_size
│        │  ├─ _select_best_pair(...)
│        │  ├─ pair_indices.pop(best_pair)
│        │  └─ _merge_pair_and_update_pair_counts(...)
│        │     ├─ _merge_pair(...)
│        │     └─ _update_pair_count(...)
│        └─ _save_result(vocab_tokens, merges)
└─ Tokenizer.save(...)
   └─ BPE.save(...)
Note

需要注意的是,当前实现与 CS336 参考实现存在一个细节差异。CS336 在 byte 层面执行 BPE,因此 frequency 相同时按照 raw bytes 的字典序进行 tie-break。而当前实现模仿 Hugging Face Tokenizers,先通过 byte-to-Unicode 映射将 bytes 转换为可逆 Unicode 字符,因此内部比较的是 Unicode 字符串的字典序。由于该映射不保证与原始 byte 顺序一致,少数情况下 merge 顺序可能与 CS336 参考实现不同,但不会影响 BPE 的整体流程。

tokenizer = dltk.Tokenizer(
    dltk.BPE(),
    pre_tokenizer=dltk.ByteLevelPreTokenizer(add_prefix_space=False),
    decoder=dltk.ByteLevelDecoder(),
    num_workers=1,
)
tokenizer.train_from_iterator(
    ['low lower lowest', 'newer wider'],
    vocab_size=300,
    special_tokens=['[UNK]', '[PAD]', '[CLS]', '[SEP]'],
    initial_alphabet=dltk.ByteLevelPreTokenizer.alphabet(),
)
print('Vocab size:', tokenizer.vocab_size)
print('First five merges:', tokenizer.model.merges[:5])
Vocab size: 277
First five merges: [('w', 'e'), ('l', 'o'), ('Ġ', 'lo'), ('Ġlo', 'we'), ('Ġlowe', 's')]

Problem (train_bpe_tinystories)

当前实验脚本从 Hugging Face roneneldan/TinyStories 读取 2,119,719 条训练样本,设置目标词表大小为 10,000,并记录训练时间和峰值内存使用。

在优化后的 BPE 实现上进行训练,结果如下:

from train_bpe_tinystories import train_bpe_tinystories

tokenizer = train_bpe_tinystories()
longest_token = max(tokenizer.get_vocab(), key=len)
print(f'Tokenizer vocab size: {tokenizer.vocab_size}')
print(f'Longest token: {longest_token!r} (length: {len(longest_token)})')
Training tokenizer on TinyStories...
Training completed in 42.7828 seconds.
[NOTE] This should be less than 30 minutes for the TinyStories dataset.
Tokenizer vocabulary size: 10000.
LRU cache info: CacheInfo(hits=5, misses=6, maxsize=100000, currsize=6)
Peak memory usage: 2.7093 GB.
[NOTE] This should be less than 30 GB for the TinyStories dataset.
Tokenizer saved to bpe_tinystories.json.
Tokenizer vocab size: 10000
Longest token: 'Ġaccomplishment' (length: 15)

根据作业要求,TinyStories 数据集训练时间应低于 30 分钟,峰值内存应低于 30 GB。当前实现的训练时间约为 16 秒,峰值内存约为 2.6 GB,满足要求。

训练完成后,tokenizer 被保存为 bpe_tinystories.json,包含 10,000 个 token 和对应的 merge rules。按照内部 byte-to-Unicode 字符串长度计算,最长 token 是 Ġaccomplishment,长度为 15。

对实验脚本进行 profiling,分析训练过程中各个阶段的时间开销:

from train_bpe_tinystories import profile_train_bpe_tinystories

profile_train_bpe_tinystories()

Profile 结果显示,当前实现主要的时间开销集中在预分词和 merge 更新阶段。

首先,regex.findall 占据了大量累计时间:

{method 'findall' of '_regex.Pattern' objects}
cumtime: 374.334s

这部分对应 ByteLevel pre-tokenization 过程。虽然单次调用开销较小,但是 TinyStories 包含超过 200 万条文本样本,大量重复调用导致累计时间较高。这说明预分词阶段仍然存在进一步优化空间,例如减少重复 token 的处理,或者通过 batch 并行减少整体等待时间。

第二个明显瓶颈是 pair frequency 的更新:

trainer.py:_merge_pair_and_update_pair_counts
cumtime: 89.856s

trainer.py:_update_pair_count
cumtime: 82.494s

这部分发生在 BPE merge loop 中。虽然当前实现已经使用 heap 快速获取最高频 pair,并通过反向索引只更新受影响的 word,但每次 merge 后仍然需要修改多个 pair 的统计信息。随着 merge 次数增加,这部分增量更新操作成为主要计算开销之一。

另外,可以看到:

utils.py:parallel_map
cumtime: 106.451s

以及 threading 相关函数占据了一定时间。这说明并行预处理本身已经引入额外的任务调度和同步开销,需要进一步关注任务粒度是否合理。如果 batch 太小,线程或进程调度成本可能抵消并行带来的收益。

同时,由于启用了多线程,cProfile 中显示的 cumtime 是所有线程累计执行时间,而不是程序实际经过的 wall-clock time。因此,profile 中某些函数的累计时间可能超过实际训练时间,这是并行执行导致的正常现象。

Problem (train_bpe_expts_owt)

2.6 BPE Tokenizer

Problem (tokenizer)

dnnlpy 状态:已在 dnnlpy.tokenizers.Tokenizer 中实现。

dnnlpy 实现支持 iterator 训练、单条和批量编码、解码、special-token 注册,以及单文件 JSON 加载和保存。相比直接返回 token id 列表,encode() 返回一个 Encoding 对象,其中包含 idstokensoffsets,方便同时获取 token 信息和原始文本位置。

与 CS336 提供的 tokenizer 接口相比,当前实现存在一些差异。当前版本没有实现 from_files(...)encode_iterable(...),而是通过单个 JSON 文件保存完整 tokenizer 状态,并在加载时将 BPE vocabulary、merge rules 和 special tokens 恢复到现有 tokenizer 对象中。

encoding = tokenizer.encode('lowest')

print('ids:', encoding.ids)
print('tokens:', encoding.tokens)
print('decoded:', tokenizer.decode(encoding.ids))
ids: [5828, 593]
tokens: ['low', 'est']
decoded: lowest

2.7 Tokenizer Experiments

Problem (tokenizer_experiments)

3.3 Basic Building Blocks

Problem (linear)

dnnlpy 状态:已在 dnnlpy.nn.Linear 中实现。

CS336 的实现不要求 bias,并在构造函数中接受 devicedtypednn.Linear 的 bias 可选但默认开启,构造函数不接收 devicedtype,并使用 kaiming_uniform 初始化而不是作业指定的 trunc_normal 初始化。用于作业实现时应传入 bias=False,再调用 .to(...)

linear = dnn.Linear(3, 2, bias=False)

x = torch.tensor([[1.0, 2.0, 3.0]])
y = linear(x)

print('weight.shape:', linear.weight.shape)
print('output.shape:', y.shape)
weight.shape: torch.Size([2, 3])
output.shape: torch.Size([1, 2])

Problem (embedding)

dnnlpy 状态:已在 dnnlpy.nn.Embedding 中实现。

与 CS336 要求的不同,dnn.Embedding 构造函数不接受 devicedtype,需要在实例化后移动或转换模块。dnnlpy 实现还提供 padding_idxmax_norm、按词频缩放梯度、预训练权重和冻结等 PyTorch 兼容参数。初始化具有作业要求的标准正态均值和方差,但没有进行截断。

embedding = dnn.Embedding(num_embeddings=8, embedding_dim=4)

token_ids = torch.tensor([[1, 3, 1]])
embedded = embedding(token_ids)

print('weight.shape:', embedding.weight.shape)
print('embedded.shape:', embedded.shape)
weight.shape: torch.Size([8, 4])
embedded.shape: torch.Size([1, 3, 4])

3.4 Transformer Block Components

Problem (rmsnorm)

dnnlpy 状态:已在 dnnlpy.nn.RMSNorm 中实现。

和 CS336 一样,实现会把 float16bfloat16 输入提升到 float32 计算归一化。eps=None 时从输入 dtype 的 machine epsilon 得到默认值,而不是固定 1e-5。它还支持多维 normalized shape 和可选 affine scaling,构造函数没有 devicedtype 参数。

rms_norm = dnn.RMSNorm(4, eps=1e-5)

x = torch.randn(2, 3, 4, dtype=torch.float16)
normalized = rms_norm(x)

print('normalized.shape:', normalized.shape)
print('normalized.dtype:', normalized.dtype)
normalized.shape: torch.Size([2, 3, 4])
normalized.dtype: torch.float32

Problem (positionwise_feedforward)

Problem (rope)

dnnlpy 状态:已在 dnnlpy.nn.RotaryPositionalEmbedding 中实现。

和 CS336 不同,dnnlpy 按需计算 sine 和 cosine,不预先缓存到 max_seq_len。因此,构造函数不需要 max_seq_len 参数。CS336 支持 token_positions 作为输入,允许指定任意 token 的位置;dnnlpy 只支持从零开始的连续位置,并且 position_offset 只能是单个整数,不能传入 tensor。

rope = dnn.RotaryPositionalEmbedding(embed_dim=4, base=10000)

x = torch.randn(2, 3, 4)
rotated = rope(x, position_offset=5)

print('rotated.shape:', rotated.shape)
rotated.shape: torch.Size([2, 3, 4])

Problem (softmax)

dnnlpy 状态:已在 dnnlpy.nn.functional.softmax 中实现。

实现会先减去指定维度上的最大值再 exponentiate,符合数值稳定版本的要求。

scores = torch.tensor([[1000.0, 1001.0, 1002.0]])
probs = dF.softmax(scores, dim=-1)

print('probs:', probs)
print('probs.sum(dim=-1):', probs.sum(dim=-1))
probs: tensor([[0.0900, 0.2447, 0.6652]])
probs.sum(dim=-1): tensor([1.])

Problem (scaled_dot_product_attention)

dnnlpy 状态:已在 dnnlpy.nn.functional.scaled_dot_product_attention 中实现。

CS336 的 boolean mask 中 True 表示该位置允许参与 attention;dnnlpy 遵循 nn.Transformer 的约定,用 True 表示该位置被屏蔽。因此传入 dnnlpy 前需要取反。函数返回 (output, weights),并额外支持 causal mask、dropout 和显式 scale。

q = torch.randn(1, 2, 3, 4)
k = torch.randn(1, 2, 3, 4)
v = torch.randn(1, 2, 3, 5)

mask = torch.ones(3, 3, dtype=torch.bool).tril()
mask = ~mask  # CS336's True -> dnnlpy's False

output, weights = dF.scaled_dot_product_attention(q, k, v, attn_mask=mask)

print('output.shape:', output.shape)
print('weights.shape:', weights.shape)
output.shape: torch.Size([1, 2, 3, 5])
weights.shape: torch.Size([1, 2, 3, 3])

Problem (multihead_self_attention)

dnnlpy 状态:已在 dnnlpy.nn.MultiheadAttention 中实现。

dnnlpy.nn.MultiheadAttention 是通用的 attention module,而不是只接收单个输入的 self-attention 接口。Bias 默认开启,可使用 dropout,并返回 (output, weights)。启用 use_rope=True 时,RoPE 使用从零开始的连续位置,不能接收用户提供的 token-position tensor。同时,boolean mask 仍采用上一题所述的反向语义。

self_attn = dnn.MultiheadAttention(
    embed_dim=8,
    num_heads=2,
    bias=False,
    use_rope=True,
)

x = torch.randn(2, 4, 8)
y, _ = self_attn(x, x, x, is_causal=True)

print('y.shape:', y.shape)
y.shape: torch.Size([2, 4, 8])

Problem (transformer_block)

Problem (transformer_lm)

3.5 Transformer Resource Accounting

Problem (transformer_accounting)

设词表大小为 \(V\)、context length 为 \(T\)、层数为 \(N\)、model dimension 为 \(D\)、attention head 数为 \(H\)、feed-forward dimension 为 \(F\)。假设 input/output embedding 不共享,且所有 linear 均无 bias,则参数量为:

\[ P = 2VD + N(4D^2 + 3DF + 2D) + D \]

对 GPT-2 XL 形状:

\[ V = 50257, \quad T = 1024, \quad N = 48, \quad D = 1600, \quad H = 25, \quad F = 4288 \]

共有 1,640,452,800 个参数。其中,float32 大约占 6.56 GB(6.11 GiB)。

单个序列的主要前向传播的 FLOPs 为:

  • Attention Q/K/V/O projections:\(N \times 8T D^2\)
  • Attention score 和 value products:\(N \times 4T^2 D\)
  • SwiGLU feed-forward projections:\(N \times 6T D F\)
  • Output LM head:\(2T D V\)

总计约 3.5168e12 FLOPs。在 \(T=1024\) 时 feed-forward 占主导;当 \(T=16384\) 时,二次方 attention 约占 61.7%,总 forward cost 约为 1.3358e14 FLOPs。

表 1:GPT-2 系列模型前向传播 FLOPs 分布
Model Attention projections Attention score/value Feed-forward LM head
Small 19.9% 13.3% 39.8% 27.1%
Medium 24.8% 12.4% 50.1% 12.7%
Large 27.3% 10.9% 54.3% 7.4%
XL 28.6% 9.2% 57.5% 4.7%

4 Training

Problem (cross_entropy)

dnnlpy 状态:已在 dnnlpy.nn.functional.cross_entropy_loss 中实现。

CS336 把最后一维视作 vocabulary,并允许任意 leading dimensions。dnnlpy 遵循 PyTorch 约定:当 dim 大于 1 时,class dimension 是第 1 维。因此当 LM logits 的 shape 为 (batch, sequence, vocab) 时,需要先展平为 (batch * sequence, vocab)。实现还支持 class weights、ignored targets、soft targets、label smoothing 和多种 reduction。

和 PyTorch 一样,dnnlpycross_entropy_loss 会在内部调用 log_softmax,因此不需要在外部显式计算 softmax。log_softmax 里调用 logsumexp 函数,使用数值稳定的方式计算 log-sum-exp。

Note

在 PyTorch 2.13 中,引入了一个新模块,nn.LinearCrossEntropyLoss。它将 LinearCrossEntropyLoss 结合在一起,采用流式处理,避免显式构造 logits tensor,从而节省内存和计算开销。详细内容请参考 PyTorch 官方文档

lm_logits = torch.randn(2, 3, 5)
targets = torch.randint(0, 5, (2, 3))

loss = dF.cross_entropy_loss(
    lm_logits.reshape(-1, 5),
    targets.reshape(-1),
)

print('loss:', loss.item())
loss: 1.8113406896591187

Problem (learning_rate_tuning)

对题目中的 quadratic SGD toy example,lr=10lr=1 更快降低 loss;lr=100 虽然激进,但由于 schedule 按 1 / sqrt(t + 1) 衰减仍可能下降;lr=1000 的早期更新 overshoot 太大,因此发散。

lr_list = [1e0, 1e1, 1e2, 1e3]
loss_list = defaultdict(list)

for lr in lr_list:
    weights = nn.Parameter(5 * torch.randn(10, 10))
    optimizer = dopt.SGD([weights], lr=lr)

    for step in range(10):
        loss = weights.pow(2).mean()
        loss.backward()
        loss_list[lr].append(loss.item())

        optimizer.step()
        optimizer.zero_grad()

df = pd.DataFrame(loss_list)
df = df.set_axis([f'lr={lr}' for lr in lr_list], axis='columns')
ipy.display(df)
lr=1.0 lr=10.0 lr=100.0 lr=1000.0
0 31.576265 27.685450 27.156178 2.950129e+01
1 30.325844 17.718687 27.156178 1.064996e+04
2 29.124941 11.339961 27.156172 3.844636e+06
3 27.971594 7.257575 27.156174 1.387914e+09
4 26.863918 4.644847 27.156172 5.010368e+11
5 25.800108 2.972703 27.156172 1.808743e+14
6 24.778423 1.902530 27.156172 6.529561e+16
7 23.797197 1.217619 27.156172 2.357172e+19
8 22.854830 0.779276 27.156170 8.509389e+21
9 21.949778 0.498737 27.156170 3.071889e+24

Problem (adamw)

dnnlpy 状态:已在 dnnlpy.optim.AdamW 中实现。

实现并维护一阶矩和二阶矩,进行 bias correction,并把 decoupled weight decay 直接应用到 parameter。它是 optim.Optimizer 的 subclass,支持 parameter groups 和 state dict。

params = nn.Parameter(torch.tensor([1.0, -1.0]))
optimizer = dopt.AdamW([params], lr=1e-3, weight_decay=0.01)

loss = params.square().sum()
loss.backward()

optimizer.step()
optimizer.zero_grad()

Problem (adamw_accounting)

设 parameter elements 为 \(P\)、需要保存的 activation elements 为 \(A\),float32 peak memory 可以近似为 \(16P + 4A\) bytes:parameters 占 \(4P\)、gradients 占 \(4P\)、两个 AdamW moments 占 \(8P\)、activations 占 \(4A\)。使用:

\[ A \approx B[N(8TD + 2HT^2 + 4TF) + TD + 2TV] \]

GPT-2 XL 中,有:

\[ P = 1,640,452,800 \qquad A \approx 4,093,347,840 \times B \]

总显存近似为:

\[ 16.373 \mathrm{GB} \times B + 26.247 \mathrm{GB} \]

按这个粗略模型,80 GB 显存最多容纳 3 个 batch。

AdamW 自身对 \(P\) 是线性复杂度,每 step 大约 \(14P\) scalar FLOPs。按题目给定的 batch 1024、400K steps、backward cost 为 forward 的两倍、有效算力为 \(0.5 \times 495\,\mathrm{TFLOP/s}\),训练约需 4,850 小时。相较于模型的 forward 和backward,optimizer 自身的 FLOPs 可忽略。

Problem (learning_rate_schedule)

dnnlpy 状态:没有实现题目要求的完整 schedule。

dnnlpy.optim.CosineAnnealingLR 提供 cosine annealing,但没有 CS336 的 linear warmup,也没有严格实现 warmup、一次 cosine decay、然后固定 minimum 的三段函数。下面是最接近的现有 scheduler。

params = nn.Parameter(torch.tensor(0.0))
optimizer = dopt.AdamW([params], lr=3e-4)
scheduler = dopt.CosineAnnealingLR(optimizer, T_max=100, eta_min=3e-5)

print('Initial learning rate:', scheduler.get_last_lr()[0])
Initial learning rate: 0.0003

Problem (gradient_clipping)

dnnlpy 状态:已在 dnnlpy.nn.utils.clip_grad_norm_ 中实现。

当使用 L2 范数时,clip_grad_norm_ 会在所有 parameter gradients 上计算 global norm,在超出阈值时原地缩放,并返回 clipping 之前的 norm。dnnlpy 还支持其他 norm orders、non-finite check 和 foreach 参数。

params = nn.Parameter(torch.tensor([3.0, 4.0]))
params.grad = torch.tensor([6.0, 8.0])

norm = dnn.utils.clip_grad_norm_([params], max_norm=1.0)

print('Before:', norm.item())
print('After:', params.grad.norm().item())
Before: 10.0
After: 0.9999999403953552

5 Training Loop

Problem (data_loading)

dnnlpy 状态:已在 dnnlpy.models.gpt.get_batch 中实现。

get_batch 随机抽取连续 window,并构造偏移一位的 next-token targets。与 CS336 不同,它接收 torch.Tensor 而不是 memory-mapped array,使用 torch.randint 采样,并把完的 batch 移动到指定 device。block_size 对应题目中的 context_length

token_stream = torch.arange(100, dtype=torch.long)
inputs, next_tokens = gpt.get_batch(
    token_stream,
    block_size=8,
    batch_size=4,
    device='cpu',
)

print('inputs.shape:', inputs.shape)
print('next_tokens.shape:', next_tokens.shape)

flag = torch.equal(inputs[:, 1:], next_tokens[:, :-1])
print('Is inputs[:, 1:] equal to next_tokens[:, :-1]?', flag)
inputs.shape: torch.Size([4, 8])
next_tokens.shape: torch.Size([4, 8])
Is inputs[:, 1:] equal to next_tokens[:, :-1]? True

Problem (checkpointing)

Problem (training_together)

6 Decoding

Problem (decoding)

7 Experiments