AI Pulse

开发者:复刻的GPT-2多3900万参数,指令微调总输原版

开发者:复刻的GPT-2多3900万参数,指令微调总输原版

当我从零开始学完如何构建LLM后,一个谜团一直困扰着我:尽管我的模型基于相同架构,却始终不如OpenAI原始的GPT-2模型。

我的模型都有1.63亿参数,设计上遵循Sebastian Raschka的《构建大语言模型(从零开始)》一书。这意味着它们与OpenAI GPT-2 "small“实例的设置几乎相同,只是我没有使用权重绑定(weight-tying),也没有在QKV矩阵上加偏置。权重绑定意味着在模型末端复用初始嵌入矩阵作为输出头,GPT-2 small使用了这一技巧,因此节省了不少参数——它的参数量是1.24亿而不是1.63亿——但至少在我自己的实验中,这以质量下降为代价。同样,虽然我发现QKV偏置在损失上能带来微小改进,但感觉那很可能在噪声范围内。

然而,在指令微调(IFT)任务上——同样改编自Raschka的书——GPT-2 small始终优于我的模型。该测试在Alpaca数据集的一个子集上微调模型,直到验证损失开始上升,然后用测试集通过训练好的模型。测试集问题的回答会被保存下来,然后我将所有被测模型的全部回答一次性交给GPT 5.5进行聚合评分;更多细节见这里。GPT-2 small在这项测试中总是比我任何一个模型表现更好。

此外,在一个更简单的评估——仅测量模型在测试集上的交叉熵损失——中,它表现得异常好。它的得分接近我自己最好的模型,并且优于其中许多模型。这个结果特别有趣的地方在于,该测试集是从我自己的训练数据中分割出来的;我的模型在训练时(至少理论上)不会看到它,但该测试集很可能与我自己模型的训练数据更相似,而不是与OpenAI的训练数据相似。

在探究这个谜团时,我检查了两件事:

- GPT-2模型很可能以现代标准来看是过度训练的;那么让我的模型也过度训练是否能让它们更接近?事实证明,不行——过度训练大概无助于IFT评估(尽管那里也许有一些信号)。不过它确实对测试损失评估帮助很大。
- 我在IFT测试中处理dropout的方式可能对某些模型有利,对另一些则不利。我决定在此评估期间统一不使用dropout,因为(对我而言反直觉的是)dropout似乎损害了大多数模型的结果,即使是那些在预训练时使用dropout的模型。尤其是OpenAI权重受到了dropout的损害,而做出一个对它们(以及我自己的一些模型)有利的改动,似乎是调查中最保守的做法。

接下来我想研究的是训练数据。

各种GPT-2模型训练所依据的确切数据集从未发布过;我们对其的了解全部来自论文,论文中说:

> [我]们创建了一个新的网络抓取,强调文档质量。为此,我们只抓取经过人工策划/过滤的网页。手动过滤整个网络抓取会非常昂贵,所以作为起点,我们抓取了Reddit(一个社交媒体平台)上所有获得至少3点karma的外链。这可以被看作是一个启发式指标,指示其他用户是否觉得该链接有趣、有教育意义或只是好玩。

他们称之为”WebText“。有一个OpenWebText试图复现它,但尽管他们尝试遵循与原始版本相同的程序,但无法保证有多大相似性。

相比之下,我通常是用FineWeb进行训练。虽然这是一个通用的网络抓取数据集,没有使用”只包含Reddit高赞帖子中链接的内容“这种策划方式,但它经过提炼,去除了明显的垃圾信息。我一直觉得两者基本等价。

但如果我错了呢?我决定看看能否通过使用更好的数据获得更好的模型。

起点

下面是我迄今为止比较过的所有模型的表格。”Test loss“列显示该模型在保留的交叉熵损失评估中的表现。”IFT epochs“列显示模型在验证损失开始上升前需要微调多少个epoch,”IFT score“是GPT 5.5对该模型在Alpaca测试集上回答给出的分数,”IFT rank“是模型在该分数上的排名。OpenAI small模型以粗体显示,我还加入了OpenAI medium模型以作对比。

模型Test lossIFT epochsIFT scoreIFT rank
OpenAI weights: medium3.231442243.751
JAX, overtrained one long epoch3.324953319.774
JAX, overtrained two normal epochs3.326482419.725
JAX, with MHA bias, no dropout3.418784418.696
JAX, no MHA bias, no dropout3.420089521.463
JAX, no MHA bias, with dropout3.476802513.2215
OpenAI weights: small3.499677226.002
1xrtx3090-stacked-interventions3.538161413.7714
8xa100m40-stacked-interventions-13.577761410.7618
Cloud FineWeb, 8x A100 40 GiB3.673623317.727
1xrtx3090-baseline3.683835415.748
8xa100m40-baseline3.691526314.1913
Cloud FineWeb, 8x H100 80 GiB3.724507414.3312
Cloud FineWeb, 8x A100 80 GiB3.729900311.3417
Cloud FineWeb, 8x B200 160 GiB3.771478414.6711
Local FineWeb train3.943522512.3116
Local FineWeb-Edu extended train4.134991515.049
Local FineWeb-Edu train4.166892514.9910

你可以看到,OpenAI small模型在测试损失方面表现相当好,考虑到它比我的模型少3900万个参数,并且测试的数据集与其可能的训练数据的差异比与我自己模型的差异更大。此外,那些胜过OpenAI small的特定模型都是用JAX而不是PyTorch训练的——我的假设是,JAX模型纯属偶然获得了更好的初始权重。

但最大的差异在于IFT分数。在本表格给出的特定运行中,OpenAI small模型得到了26.00分——我自己的模型中最好的也只有21.46分,低了4.5分以上。

这种差异在我所有其他测试运行中都是一致的。GPT-2 small模型总是领先于我的模型。(当然,GPT-2 medium超过了GPT-2 small和我所有的模型,但鉴于它的规模是我的两倍,这并不令人惊讶。)

其实,很久以前我就曾尝试将数据质量作为提升模型性能的杠杆。在表格底部,你可以看到两个测试损失最差的模型:

- ”Local FineWeb-Edu train"
- "Local FineWeb-Edu extended train“

这两个(正如你可能从名字猜到的)是在FineWeb-Edu数据集上训练的,该数据集只包含FineWeb中最”教育性“的数据。它们在测试损失上得分非常差。鉴于测试数据集来自FineWeb,这并不奇怪——正如我之前写过的:

> 如果你用一个在简·奥斯汀上训练的模型来评估查克·廷格尔,你不会得到令人惊奇的结果。

但同样,GPT-2也有同样的问题,却在测试损失评估中表现得非常好。

另一方面,虽然这些FineWeb-Edu模型在IFT评估上的表现并不出色——还有很多我的其他模型排在它们前面——但它们似乎确实有点超常发挥。在我做过的所有IFT评估中,它们一直比其他许多模型得分更高——尽管它们在测试评估上的损失很差。

此外:它们是我最早训练的那批模型,那是在我花时间学习如何优化超参数和训练循环之前。它们没有使用梯度裁剪,使用了dropout,批次大小只是”我能塞进GPU的量“,而且我没有将学习率设置到合适的值,也没有在训练过程中安排学习率计划。

所以,也许在新的FineWeb-Edu上使用我的训练改进重新训练会有帮助?也许对训练数据做其他调整也值得研究?

计划

我决定看看如果用更高质量的数据训练模型会发生什么。具体来说,我会使用我当前优化的循环和超参数,在四个不同的数据集上训练模型:

- FineWeb-Edu——本质上与”Local FineWeb-Edu train“相同,但训练设置更好。这将检验”更多教育性 -> 更好“的假设。
- FineWeb和FineWeb-Edu的50:50混合。我读到过,LLM在训练数据中包含一定量的低质量数据可以有帮助,因为这有助于它们泛化。也许在FineWeb-Edu之外再加一些FineWeb会改善测试损失,同时也对IFT测试有帮助?
- 一个”策划“数据集,包含45%来自FineWeb的内容、45%来自FineWeb-Edu、10%来自Simple English Wikipedia。完整的维基百科非常庞大,充满了冷门知识——而Simple English版本很小,希望每个token包含更有用的信息。而且方便的是,Answer.ai已经在Hugging Face Hub上提供了一个快照。刻意在训练集中放入大量百科全书式数据,会不会让模型在IFT评估中表现更好(那里有很多事实性问题,比如”谁写了《傲慢与偏见》“)?
- OpenWebText。尽管我不确定它与原始WebText的匹配程度,但既然它就在那里,不去尝试在上面训练并看看效果似乎很愚蠢。

我会在所选数据集的32亿个token上训练每个模型;这是对我163M参数模型的Chinchilla最优量。如果有什么有趣的结果,那么我可能会考虑以后做过度训练的模型。

我决定至少要有点科学态度,并预先注册一些预测:

- 纯FineWeb-Edu模型在测试损失上会表现很差,但比我旧的FineWeb-Edu模型更好(90%置信度)。它也会在IFT评估中超常发挥(90%置信度)。
- 50:50混合:我预计它在测试评估上会比我的JAX纯FineWeb模型差(70%),但比FineWeb-Edu模型好(90%)。我不确定它在IFT评估上会怎样,但认为可能介于两组之间(60%)。
- 策划数据集:我对它在IFT评估上寄予厚望——说80%的几率它会是我所有模型中最好的。对于测试损失评估,我预计它的表现会和50:50混合差不多,也许稍微差一点(70%)。
- 我不知道OpenWebText评估会怎样!可能更差,可能更好。

结果是这样的。

FineWeb-Edu模型

我已经有一个基于FineWeb-Edu的数据集了,当时训练那两个原始模型时准备过。那只是我去年12月生成时原始数据集的100亿token样本,格式适合我的训练脚本(详见数据集卡)。

我用JAX代码(这个系列的其他帖子一直在用)启动了一次训练运行:

giles@poppy:~/Dev/jax-gpt2-from-scratch (main)$ XLA_PYTHON_CLIENT_MEM_FRACTION=0.95 uv run train.py full-llm-full-train-with-mha-output-bias-fineweb-edu datasets/
2026-09-11 18:11:47.991583 Downloading dataset
Fetching 4 files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 4/4 [00:00<00:00, 1772.93it/s]
Download complete: : 0.00B [00:00, ?B/s]                                                                                                            | 0/4 [00:00<?, ?it/s]
2026-09-11 18:11:48.226273 Loading dataset into RAM
Download complete: : 0.00B [00:00, ?B/s]
2026-09-11 18:16:29.507646 Creating model
2026-09-11 18:16:33.042509 Creating optimizer
2026-09-11 18:16:34.138990 Start train
  0%|                                                                                                                                           | 0/33165 [00:00<?, ?it/s]
2026-09-11 18:17:38.486288 Saving checkpoint
  1%|▌                                                                                                     | 173/33165 [13:22<39:17:03,  4.29s/it, loss=6.897, tps=21,201]

……不到40小时后,我得到了一个模型:

Training complete in 142,912.226 seconds
2026-09-13 09:58:26.437276 Tokens seen: 3,260,252,160
2026-09-13 09:58:26.437284 Throughput: 22,813 tokens/second
2026-09-13 09:58:26.437302 Final train loss: 3.342
2026-09-13 09:58:26.437309 Done

我将保存的JAX safetensors文件中最后一个检查点转换成了与我的PyTorch评估代码兼容的格式,并运行了冒烟测试:它会如何补全句子”Every effort moves you"?

> Every effort moves you closer to God“s Kingdom, and even closer to Him. As we can see in

那很好,也很连贯——虽然异常宗教化!——所以很有希望。我运行了测试评估:

giles@perry:~/Dev/ddp-base-model-from-scratch (main)$ uv run test_loss.py datasets/ ../jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-fineweb-edu/model.json ../jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-fineweb-edu/checkpoints/latest/pytorch-model.safetensors
Fetching 4 files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 4/4 [00:00<00:00, 2758.50it/s]
100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 3200/3200 [03:52<00:00, 13.74it/s]
Loss against our test dataset: 3.632900

那相当不错,测试损失比我所有在没有优化超参数的情况下训练的模型都要好,而比我所有在FineWeb上使用优化超参数训练的模型都要差。所以这符合我的预测,即它会比旧的FineWeb-Edu模型好;它居然也比那些未优化的FineWeb训练运行更好,这一点似乎很合理,以至于我觉得自己没预测到它会恰好落到那个位置真是愚蠢:-)

我决定把IFT评估留到最后,这样我就可以一起检查这些实验中的所有模型,所以是时候把这个上传到Hugging Face,然后继续下一个模型了。

FineWeb与FineWeb-Edu的50:50混合

我建立了一个新仓库,里面有一个专门为我的训练设置准备数据集的脚本。你提供配置来指定一些源数据集以及如何处理它们和如何混合它们,它就会将具有所需特性的新数据集上传到Hugging Face Hub。

例如,对于FineWeb与FineWeb-Edu的50:50分割,配置如下:

{
    "seed": 42,
    "tokens_desired": 10000000000,
    "upload_dataset_name": "gpjt/fw-fwedu-5050-gpt2-tokens",
    "sources": [
        {
            "name": "FineWeb",
            "hf_id": "HuggingFaceFW/fineweb",
            "hf_name": "sample-10BT",
            "hf_split": "train",
            "item_field": "text",
            "weight": 50
        },
        {
            "name": "FineWeb-Edu",
            "hf_id": "HuggingFaceFW/fineweb-edu",
            "hf_name": "sample-10BT",
            "hf_split": "train",
            "item_field": "text",
            "weight": 50
        }
    ]
}

脚本的工作方式非常简单:它根据这些权重和tokens_desired计算出它想从每个源数据集中获取多少token,打乱源数据集中的条目,然后循环直到输出中存储了所需数量的token或更多。在循环中,它计算出哪个源当前最欠缺,从中抓取一个条目,进行token化,并将其添加到输出中。

用这个50:50配置运行似乎效果不错:

giles@perry:~/Dev/prepare-llm-training-dataset (main)$ uv run prepare-dataset.py runs/fw-fwedu-5050/
Resolving data files: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████| 27468/27468 [00:00<00:00, 89875.56it/s]
Loading dataset shards: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████| 102/102 [00:00<00:00, 133.75it/s]
Resolving data files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 2410/2410 [00:00<00:00, 87461.48it/s]
Loading dataset shards: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████| 98/98 [00:00<00:00, 200.09it/s]
2026-09-13 20:13:22.000187: Generating dataset; per-source counts
2026-09-13 20:13:22.000217: FineWeb: 5,000,000,000
2026-09-13 20:13:22.000221: FineWeb-Edu: 5,000,000,000
FineWeb: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 4999999705/5000000000 [1:01:33<00:00, 1353639.33token/s]
FineWeb-Edu: 5000000363token [1:01:33, 1353639.47token/s]
2026-09-13 21:14:55.747239:

Done generating tokens
2026-09-13 21:14:55.748480: FineWeb: 4,999,999,705 / 5,000,000,000 (1.000, 1 iterators)
2026-09-13 21:14:55.748487: FineWeb-Edu: 5,000,000,363 / 5,000,000,000 (1.000, 1 iterators)
2026-09-13 21:14:55.748489: Total: 10,000,000,068
2026-09-13 21:14:55.748491: Catting...
2026-09-13 21:16:29.565152: Catted into a tensor of shape torch.Size([10000000068])
2026-09-13 21:16:29.566663: Saving...
2026-09-13 21:16:36.006267: Saved
2026-09-13 21:16:36.009413: Uploading to gpjt/fw-fwedu-5050-gpt2-tokens
Processing Files (1 / 1)      : 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 20.0GB / 20.0GB,  117MB/s
New Data Upload               : 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 14.6GB / 14.6GB, 98.1MB/s
  ...du-5050/train.safetensors: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 20.0GB / 20.0GB
2026-09-13 21:17:59.545875: Done

所以我们几乎完美地实现了数据集之间50:50的平衡,并将这个数据集保存在了Hugging Face上。

我运行了一个脚本来仔细检查它看起来是否合理,结果确实如此,所以是时候启动训练运行了:

giles@perry:~/Dev/jax-gpt2-from-scratch (main)$ XLA_PYTHON_CLIENT_MEM_FRACTION=0.90 uv run train.py full-llm-full-train-with-mha-output-bias-fw-fwedu-5050 datasets/
2026-09-13 21:20:59.880918 Downloading dataset
Fetching 2 files: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 2/2 [01:13<00:00, 36.70s/it]
Download complete: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████| 20.0G/20.0G [01:13<00:00, 1.24GB/s]
2026-09-13 21:22:13.521745 Loading dataset into RAM
Download complete: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████████| 20.0G/20.0G [01:13<00:00, 272MB/s]
2026-09-13 21:22:33.787720 Creating model
2026-09-13 21:22:35.501063 Creating optimizer
2026-09-13 21:22:36.043837 Start train
  0%|                                                                                                                                           | 0/33165 [00:00<?, ?it/s]
2026-09-13 21:23:11.437206 Saving checkpoint
  0%|                                                                                                       | 26/33165 [02:20<38:07:05,  4.14s/it, loss=9.308, tps=18,246]

那是在perry上运行的,我的普通工作站,我与下面在poppy(我的训练机)上的”策划“模型训练运行并行启动,但为了这个帖子的目的,我会把运行分开叙述。

运行了一个小时左右后,我们断电了。我猜是烘干机开着、汽车充电、水壶烧水、电磁炉开着,再加上两台机器做训练运行,对咱们家的电力来说有点太多了……这在未来可能会是个问题,尤其是如果(按计划)我把poppy变成多GPU机器的话。

不过,就目前而言,我在重新合上断路器后能够再次启动它,而且情况一直稳定。

又是大约40小时后:

Training complete in 136,060.457 seconds
2026-09-15 12:05:26.432638 Tokens seen: 3,227,516,928
2026-09-15 12:05:26.432642 Throughput: 23,721 tokens/second
2026-09-15 12:05:26.432650 Final train loss: 3.793
2026-09-15 12:05:26.432653 Done

(注意,像这样重新启动的运行结束时报告的数字只包括重启后发生的事情。)

我将其转换为PyTorch兼容的张量,并做了冒烟测试:

> Every effort moves you on to other options—in fact, it”s not even worth that effort. Just make

看起来不错!是时候进行损失测试了:

giles@perry:~/Dev/ddp-base-model-from-scratch (main)$ uv run test_loss.py datasets/ ../jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-fw-fwedu-5050/model.json ../jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-fw-fwedu-5050/checkpoints/latest/pytorch-model.safetensors
Fetching 4 files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 4/4 [00:00<00:00, 1192.07it/s]
100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 3200/3200 [03:53<00:00, 13.72it/s]
Loss against our test dataset: 3.462454

这几乎符合我的预测,它会比JAX纯FineWeb模型差,只是它比其中最差的"JAX, no MHA bias, with dropout“要好:实际上它比我预测的还要好。

所以,这是一个有希望的模型。是时候上传到Hugging Face了——现在继续下一个。

”策划“数据集

用我的数据集准备脚本,这很容易设置:

{
    "seed": 42,
    "tokens_desired": 10000000000,
    "upload_dataset_name": "gpjt/fw-fwedu-simplewiki-gpt2-tokens",
    "sources": [
        {
            "name": "FineWeb",
            "hf_id": "HuggingFaceFW/fineweb",
            "hf_name": "sample-10BT",
            "hf_split": "train",
            "item_field": "text",
            "weight": 45
        },
        {
            "name": "FineWeb-Edu",
            "hf_id": "HuggingFaceFW/fineweb-edu",
            "hf_name": "sample-10BT",
            "hf_split": "train",
            "item_field": "text",
            "weight": 45
        },
        {
            "name": "Simple English Wikipedia",
            "hf_id": "answerdotai/simplewiki",
            "hf_name": "articles",
            "hf_split": "train",
            "item_field": "md",
            "weight": 10
        }
    ]
}

运行起来很顺利:

giles@perry:~/Dev/prepare-llm-training-dataset (main)$ uv run prepare-dataset.py runs/fw-fwedu-simplewiki/
Resolving data files: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████| 27468/27468 [00:00<00:00, 90196.13it/s]
Loading dataset shards: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████| 102/102 [00:00<00:00, 358.90it/s]
Resolving data files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 2410/2410 [00:00<00:00, 88254.11it/s]
Loading dataset shards: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████| 98/98 [00:00<00:00, 589.23it/s]
2026-09-13 18:59:04.106327: Generating dataset; per-source counts
2026-09-13 18:59:04.106387: FineWeb: 4,500,000,000
2026-09-13 18:59:04.106407: FineWeb-Edu: 4,500,000,000
2026-09-13 18:59:04.106422: Simple English Wikipedia: 1,000,000,000
FineWeb: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 4499997964/4500000000 [59:41<00:00, 1256362.56token/s]
FineWeb-Edu: 4500000607token [59:41, 1256363.31token/s]
Simple English Wikipedia: 1000002889token [59:41, 279192.58token/s]
2026-09-13 19:58:45.874744:

Done generating tokens
2026-09-13 19:58:45.876043: FineWeb: 4,499,997,964 / 4,500,000,000 (1.000, 1 iterators)
2026-09-13 19:58:45.876048: FineWeb-Edu: 4,500,000,607 / 4,500,000,000 (1.000, 1 iterators)
2026-09-13 19:58:45.876052: Simple English Wikipedia: 1,000,002,889 / 1,000,000,000 (1.000, 6 iterators)
2026-09-13 19:58:45.876054: Total: 10,000,001,460
2026-09-13 19:58:45.876056: Catting...
2026-09-13 20:00:18.811748: Catted into a tensor of shape torch.Size([10000001460])
2026-09-13 20:00:18.813169: Saving...
2026-09-13 20:00:22.773873: Saved
2026-09-13 20:00:22.773936: Uploading to gpjt/fw-fwedu-simplewiki-gpt2-tokens
Processing Files (1 / 1)      : 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 20.0GB / 20.0GB,  143MB/s
New Data Upload               : 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 19.8GB / 19.8GB,  142MB/s
  ...plewiki/train.safetensors: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 20.0GB / 20.0GB
2026-09-13 20:01:59.270021: Done

在那个输出中值得注意的是Simple English Wikipedia的”6 iterators“。如果源数据集在这个脚本构建结果时用完了条目,我们会用不同的shuffle种子重新开始迭代它,这样顺序会不同。”6 iterators“意味着它需要这样做6次——脚本开始时创建迭代器算一次,然后还有五次。所以这意味着Simple English Wikipedia在数据集中重复(过采样)了大约五到六次。

这不是坏事!根据我读到的内容,在LLM训练数据集中对高教育性内容进行过采样实际上是很标准的。而且,脚本生成的数据集是100亿token,我们在这个帖子中只使用其中32亿用于训练运行,所以它只会出现大约一到两次。重复只有在(如果)我们对数据集进行过度训练时才会真正起作用。

无论如何,我检查了一下上传的数据集——前几个条目显然来自FineWeb、FineWeb-Edu和Simple English Wikipedia。

是时候启动一次训练运行了:

giles@poppy:~/Dev/jax-gpt2-from-scratch (main)$ XLA_PYTHON_CLIENT_MEM_FRACTION=0.95 uv run train.py full-llm-full-train-with-mha-output-bias-fw-fwedu-simplewiki datasets/
2026-09-13 20:24:48.037024 Downloading dataset
Downloading (incomplete total...): 0.00B [00:00, ?B/s]                                                                                                                   Warning: You are sending unauthenticated requests to the HF Hub. Please set a HF_TOKEN to enable higher rate limits and faster downloads.            | 0/2 [00:00<?, ?it/s]
WARNING:huggingface_hub.utils._http:Warning: You are sending unauthenticated requests to the HF Hub. Please set a HF_TOKEN to enable higher rate limits and faster downloads.
Fetching 2 files: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 2/2 [02:51<00:00, 85.85s/it]
Download complete: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████████| 20.0G/20.0G [02:51<00:00, 435MB/s]
2026-09-13 20:27:39.934884 Loading dataset into RAM
Download complete: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████████| 20.0G/20.0G [02:51<00:00, 116MB/s]
2026-09-13 20:31:20.492877 Creating model
2026-09-13 20:31:24.054143 Creating optimizer
2026-09-13 20:31:25.100832 Start train
  0%|                                                                                                                                           | 0/33165 [00:00<?, ?it/s]
2026-09-13 20:32:29.650379 Saving checkpoint
  0%|▎                                                                                                     | 107/33165 [08:38<39:05:39,  4.26s/it, loss=7.631, tps=20,293]

同样,它被影响50:50训练运行的那次断电打断了,但我能够从检查点重新启动。

又过了22小时后,它崩溃了,出现了一个我以前见过的错误:

jax.errors.JaxRuntimeError: INTERNAL: CUDA error: Failed to end stream capture: CUDA_ERROR_STREAM_CAPTURE_INVALIDATED: operation failed due to a previous error during capture [executable_name='jit_train_step']

上次遇到时我将其视为一次性怪事,但这次我深入调查了一下。我注意到它从未在perry上发生过,但似乎是poppy上的问题,而poppy有较旧版本的CUDA和Nvidia驱动——这可能是原因吗?我决定在启动下一次运行之前升级它们,但暂时只是从最近的检查点重新启动了运行。(注意:对于遇到同样错误的任何人:升级后就再也没有发生过,所以值得一试。)

这次它顺利完成了:

Training complete in 59,564.515 seconds
2026-09-15 15:56:52.909888 Tokens seen: 1,367,212,032
2026-09-15 15:56:52.909894 Throughput: 22,953 tokens/second
2026-09-15 15:56:52.909912 Final train loss: 3.332
2026-09-15 15:56:52.909959 Done

同样,这些数字只显示了最近一次重启后发生的事情。

我把它复制到perry,转换成与我的PyTorch代码兼容的格式,并运行了冒烟测试:

> Every effort moves you by the air, for it will make you a better athlete, so your body becomes bigger and stronger

足够连贯——是时候进行损失评估了:

giles@perry:~/Dev/ddp-base-model-from-scratch (main)$ uv run test_loss.py datasets/ ~/Dev/jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-fw-fwedu-simplewiki/model.json ~/Dev/jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-fw-fwedu-simplewiki/checkpoints/latest/pytorch-model.safetensors
Fetching 4 files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 4/4 [00:00<00:00, 1007.64it/s]
100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 3200/3200 [03:57<00:00, 13.48it/s]
Loss against our test dataset: 3.542460

再次符合我的预测——比JAX纯FineWeb模型差,实际上也比最好的PyTorch模型1xrtx3090-stacked-interventions差,也比50:50混合差,但比FineWeb-Edu模型好。

我把它上传到了Hugging Face,然后该继续这组实验本应是的最后一个模型了。

OpenWebText运行

同样,这是一个足够简单的配置:

{
    "seed": 42,
    "tokens_desired": 10000000000,
    "upload_dataset_name": "gpjt/openwebtext-gpt2-tokens",
    "sources": [
        {
            "name": "OpenWebText",
            "hf_id": "Skylion007/openwebtext",
            "hf_name": "plain_text",
            "hf_split": "train",
            "item_field": "text",
            "weight": 50
        }
    ]
}

……构建和上传过程很顺利(而且花的时间少得多——出于某种原因,从单个数据集中随机抽样比从两三个数据集中抽样更快):

giles@perry:~/Dev/prepare-llm-training-dataset (main)$ uv run prepare-dataset.py runs/openwebtext/
Resolving data files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████| 80/80 [00:00<00:00, 32723.26it/s]
Resolving data files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████| 80/80 [00:00<00:00, 97940.55it/s]
Loading dataset shards: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████| 80/80 [00:00<00:00, 1200.13it/s]
2026-09-15 13:16:47.622617: Generating dataset; per-source counts
2026-09-15 13:16:47.622645: OpenWebText: 10,000,000,000
Resolving data files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████| 80/80 [00:00<00:00, 45602.65it/s]
Resolving data files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████| 80/80 [00:00<00:00, 67650.06it/s]
Loading dataset shards: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████| 80/80 [00:00<00:00, 307.11it/s]
OpenWebText: 10000000024token [31:46, 5246208.64token/s]
2026-09-15 13:48:33.761350:

Done generating tokens
2026-09-15 13:48:33.762021: OpenWebText: 10,000,000,024 / 10,000,000,000 (1.000, 2 iterators)
2026-09-15 13:48:33.762026: Total: 10,000,000,024
2026-09-15 13:48:33.762028: Catting...
2026-09-15 13:49:33.115508: Catted into a tensor of shape torch.Size([10000000024])
2026-09-15 13:49:33.115923: Saving...
2026-09-15 13:49:36.365978: Saved
2026-09-15 13:49:36.366027: Uploading to gpjt/openwebtext-gpt2-tokens
Processing Files (0 / 1)      : 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████▉| 20.0GB / 20.0GB,  147MB/s
New Data Upload               : 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 19.9GB / 19.9GB,  147MB/s
  ...webtext/train.safetensors: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████▉| 20.0GB / 20.0GB
2026-09-15 13:51:16.202890: Done

注意它需要过采样——那个”2 iterators“。OpenWebText未压缩大约40GiB,所以大约100亿个GPT-2 token——大概稍微少一点。同样,鉴于我计划只使用数据集的前32亿token,我觉得这不会有影响。

我在新上传的Hugging Face数据集上运行了检查脚本,一切正常,所以训练运行准备就绪。

我先用sudo pacman -Syu升级了poppy,看看是否能解决上次运行中遇到的奇怪错误(正如我所说,看起来确实解决了),然后启动它:

giles@poppy:~/Dev/jax-gpt2-from-scratch (main)$ XLA_PYTHON_CLIENT_MEM_FRACTION=0.95 uv run train.py full-llm-full-train-with-mha-output-bias-openwebtext datasets/
2026-09-15 16:42:32.606185 Downloading dataset
Fetching 2 files: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 2/2 [00:00<00:00, 941.38it/s]
Download complete: : 0.00B [00:00, ?B/s]                                                                                                            | 0/2 [00:00<?, ?it/s]
2026-09-15 16:42:32.879987 Loading dataset into RAM
Download complete: : 0.00B [00:00, ?B/s]
2026-09-15 16:45:40.438791 Creating model
2026-09-15 16:45:43.840269 Creating optimizer
2026-09-15 16:45:44.848351 Start train
  0%|                                                                                                                                           | 0/33165 [00:00<?, ?it/s]
2026-09-15 16:46:50.632075 Saving checkpoint
  1%|█                                                                                                     | 332/33165 [24:33<38:45:54,  4.25s/it, loss=6.623, tps=22,154]

大约31小时后,它又崩溃了,但这次是我自己的愚蠢错误:poppy的磁盘相对较小,我空间不够了。我修复了这个问题,从最近的检查点重新启动,这次它完成了:

Training complete in 33,927.995 seconds
2026-09-17 11:25:10.835989 Tokens seen: 779,747,328
2026-09-17 11:25:10.835994 Throughput: 22,982 tokens/second
2026-09-17 11:25:10.836012 Final train loss: 3.165
2026-09-17 11:25:10.836018 Done

我将其转换为PyTorch进行冒烟测试:

> Every effort moves you through each phase, so it”s not a complete picture. I“m sure your story was

……看起来不错,所以是时候进行测试损失评估了:

giles@perry:~/Dev/ddp-base-model-from-scratch (main)$ uv run test_loss.py datasets/ ~/Dev/jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-openwebtext/model.json ~/Dev/jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-openwebtext/checkpoints/latest/pytorch-model.safetensors
Fetching 4 files: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 4/4 [00:00<00:00, 674.76it/s]
100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 3200/3200 [03:59<00:00, 13.37it/s]
Loss against our test dataset: 4.045255

这是我们这个实验中最差的分数!比我至今所有模型都差,除了那两个没有优化超参数的FineWeb-Edu模型。

不过,这个帖子的初稿直接从这里开始写结果,但故事还没结束……

测试集污染

GPT-6 Astra毫不留情。在我发布这些帖子中的任何一篇之前,我都会让一组LLM编辑委员会检查问题。GPT-6 Astra不仅检查了文本,还访问了我链接的代码来检查,并发现了一个问题。

事后看来很明显,但我构建新数据集的代码有很高风险包含(理论上保留的)测试集的内容。测试集的生成方式是去年12月我下载了FineWeb的100亿样本,将其分成99%的训练数据和1%的”验证“数据。该验证分割大约有1亿token,而我实际训练期间只用了其中前1900万左右进行验证运行,所以我(有点武断地)指定了其中从位置5000万开始的另外约1900万token作为我的测试集。

现在,我的新数据集生成代码只是从完整的100亿FineWeb样本中随机抽样。所以没有什么能阻止它拉入旧验证分割中的数据!这意味着我的新”策划“和”50:50“数据集很可能至少包含了一些本应在训练期间保留给模型的测试集。

反思起来,问题可能更严重。FineWeb-Edu是FineWeb的子集;我现有的FineWeb-Edu数据集来自Hugging Face原始数据集的100亿样本,因此它也可能包含我放入测试集中的文档。

首先要做的是确定问题的规模。我写了一个脚本来接收一个”禁止“数据集和分割;假设其格式为一个大的GPT-2 token张量,这是我所有数据集的样子。然后它会按文本结束token进行分割,并为每个生成的”文档“生成哈希和token计数。可选地,你可以将其限制为只考虑一个子集——从位置p开始的n个token——然后它会为该切片内的文档生成哈希/长度,或者与切片在开头或结尾重叠的文档。

我运行它生成了整个验证集的哈希列表——gpjt/fineweb-gpt2-tokens的验证分割——然后用第二个脚本检查我的各个训练集(以及验证集本身),看看污染问题有多严重。我得到以下结果:

数据集分割与验证集的污染
gpjt/fineweb-gpt2-tokensvalidation102163003 / 102163003 tokens (100.00%)
gpjt/fineweb-gpt2-tokenstrain636166 / 102163003 tokens (0.62%)
gpjt/fineweb-edu-gpt2-tokenstrain672189 / 102163003 tokens (0.66%)
gpjt/fw-fwedu-5050-gpt2-tokenstrain49224580 / 102163003 tokens (48.18%)
gpjt/fw-fwedu-simplewiki-gpt2-tokenstrain44233824 / 102163003 tokens (43.30%)
gpjt/openwebtext-gpt2-tokenstrain212 / 102163003 tokens (0.00%)

所以:

- 验证集与自身100%”污染“,这是一个有用的健全性检查。
- gpjt/fineweb-gpt2-tokens的训练集有我感觉到的小量污染。有趣的是竟然有任何污染——我认为这一定意味着原始数据集中有一些重复文档,其中一些最终副本同时出现在我的训练和验证分割中。
- gpjt/fineweb-edu-gpt2-tokens数据集也有一个让我感到安心的低水平污染。
- 然而,gpjt/fw-fwedu-5050-gpt2-tokens和gpjt/fw-fwedu-simplewiki-gpt2-tokens看起来都有问题。在这两种情况下,训练数据集中都包含了超过40%的验证/测试集。
- gpjt/openwebtext-gpt2-tokens如你所料几乎完全没有污染。看起来可能有一个文档恰好被OpenWebText和FineWeb爬虫都抓到了,然后被包含在我用于验证的FineWeb部分中。

然而,这些数字——虽然吓人,至少对50:50和策划数据集而言——并不是应该使用的数字。它们显示了整个验证集在完整训练集中出现了多少;我真正关心的是测试集——验证分割中从位置5000万开始的1900万token——中有多少出现在我实际训练所用的训练数据集子集中——即它们的前约32亿token中。

我重新运行脚本只为测试集生成哈希,然后重新运行污染检查脚本,告诉它只看训练token的适当子集,得到:

数据集(仅前32亿token)分割与测试集的污染
gpjt/fineweb-gpt2-tokenstrain26557 / 19632681 tokens (0.14%)
gpjt/fineweb-edu-gpt2-tokenstrain32079 / 19632681 tokens (0.16%)
gpjt/fw-fwedu-5050-gpt2-tokenstrain2986889 / 19632681 tokens (15.21%)
gpjt/fw-fwedu-simplewiki-gpt2-tokenstrain2682430 / 19632681 tokens (13.66%)
gpjt/openwebtext-gpt2-tokenstrainNone

很明显存在问题——肯定是对gpjt/fw-fwedu-5050-gpt2-tokens和gpjt/fw-fwedu-simplewiki-gpt2-tokens。它们训练时看到了感觉相当数量的测试集,所以它们在测试损失评估上的结果充其量是可疑的。

我决定重新训练那两个模型,看看损失方面的结果。如果差异很大,我会调查gpjt/fineweb-gpt2-tokens和gpjt/fineweb-edu-gpt2-tokens(小得多的)污染风险。但如果很小,我不会太担心。

我扩展了准备数据集的脚本,使配置文件可以指定forbidden_dataset。源数据集中与禁止文档匹配的任何文档都将从输出中排除。然后我更新了gpjt/fw-fwedu-5050-gpt2-tokens和gpjt/fw-fwedu-simplewiki-gpt2-tokens的配置,使gpjt/fineweb-gpt2-tokens的整个验证分割被禁止,并重新生成它们。你可以在这里和这里看到更新的数据集。对它们运行污染检查脚本显示它们干净了。

然后我重新做了那些模型的完整训练运行;无污染的50:50分割模型在这里,策划模型在这里。

好消息是:两者实际上在测试损失评估上都比它们受污染数据训练的对应版本做得稍微好一点:

模型受污染Test loss
JAX, FineWeb/FineWeb-Edu 50:50No3.449257
JAX, FineWeb/FineWeb-Edu 50:50Yes3.462454
JAX, curatedNo3.534068
JAX, curatedYes3.542460

想到几种可能性;也许在像这样小的163M模型上,从测试集学习根本不会发生,或者也许虽然受污染模型在学习时,它们从中获得的好处被那些替代测试集数据的数据抵消了,那些数据在训练目的上可能在某种程度上更好,至少在损失评估方面是如此。

但无论如何,我觉得如果在训练期间看到超过10%的测试集数据的效果如此微小,那么看到不到0.2%——这组训练运行中的FineWeb-Edu模型所看到的量,以及之前实验中我所有其他纯FineWeb模型所看到的量——的效果会更小,我会忽略它。

那是个极好的消息!我不需要从零开始所有的实验。

在本文的其余部分,我将同时包含受污染模型和无污染模型的数字和结果——它们因多种原因很有趣——但在未来的帖子中,我会跳过受污染的模型。

所以——终于!——让我们开始深入最终结果。

结果

首先,我认为值得在上下文中查看所有测试损失结果。它们在下面的表格中,新模型用粗体显示:

模型Test loss
OpenAI weights: medium3.231442
JAX, overtrained one long epoch3.324953
JAX, overtrained two normal epochs3.326482
JAX, with MHA bias, no dropout3.418784
JAX, no MHA bias, no dropout3.420089
JAX, FineWeb/FineWeb-Edu 50:50 (uncontaminated)3.449257
JAX, FineWeb/FineWeb-Edu 50:50 (contaminated)3.462454
JAX, no MHA bias, with dropout3.476802
OpenAI weights: small3.499677
JAX, curated (uncontaminated)3.534068
1xrtx3090-stacked-interventions3.538161
JAX, curated (contaminated)3.542460
8xa100m40-stacked-interventions-13.577761
JAX, FineWeb-Edu3.632900
Cloud FineWeb, 8x A100 40 GiB3.673623
1xrtx3090-baseline3.683835
8xa100m40-baseline3.691526
Cloud FineWeb, 8x H100 80 GiB3.724507
Cloud FineWeb, 8x A100 80 GiB3.729900
Cloud FineWeb, 8x B200 160 GiB3.771478
Local FineWeb train3.943522
JAX, openwebtext4.045255
Local FineWeb-Edu extended train4.134991
Local FineWeb-Edu train4.166892

我认为这里有一个非常明确的信息:对于新模型,训练混合中FineWeb越多,模型在这个评估上就越好。我想在运行这些实验之前做预测时,我可能潜意识里已经预料到了,但事后看来这是如此极其明显,以至于我觉得没有明确提出来真是愚蠢!

但这告诉了我们一些有趣的事情。从论文中的描述来看,无论OpenAI在什么数据上做GPT-2训练运行,它都不像FineWeb。它可能更接近OpenWebText——然而,那个模型是测试评估中表现最差的,所以如果它更像OpenWebText,那么一定涉及其他因素。

但现在继续:IFT测试怎么样——那个最初引发这一切工作的测试?

我为所有新模型生成了一组IFT响应,然后将它们(以及上表中所有其他模型的响应)交给GPT 5.5,发现我的一个新模型已经非常接近原始GPT-2 small权重了!所以我做了四次额外运行,以便得到平均值。

以下是结果——”IFT score“是判断器所有五次运行的平均值,”IFT rank“基于此排名。”IFT epochs“来自原始结果生成脚本。

模型Test lossIFT epochsIFT scoreIFT rank
OpenAI weights: medium3.231442242.361
JAX, overtrained one long epoch3.324953318.677
JAX, overtrained two normal epochs3.326482418.716
JAX, with MHA bias, no dropout3.418784417.908
JAX, no MHA bias, no dropout3.420089520.504
JAX, FineWeb/FineWeb-Edu 50:50 (uncontaminated)3.449257417.699
JAX, FineWeb/FineWeb-Edu 50:50 (contaminated)3.462454419.305
JAX, no MHA bias, with dropout3.476802513.0221
OpenAI weights: small3.499677225.192
JAX, curated (uncontaminated)3.534068416.6310
1xrtx3090-stacked-interventions3.538161413.5119
JAX, curated (contaminated)3.542460413.5818
8xa100m40-stacked-interventions-13.577761410.1924
JAX, FineWeb-Edu3.632900424.563
Cloud FineWeb, 8x A100 40 GiB3.673623316.5911
1xrtx3090-baseline3.683835415.1512
8xa100m40-baseline3.691526313.6416
Cloud FineWeb, 8x H100 80 GiB3.724507413.5917
Cloud FineWeb, 8x A100 80 GiB3.729900310.7923
Cloud FineWeb, 8x B200 160 GiB3.771478413.7015
Local FineWeb train3.943522511.8722
JAX, openwebtext4.045255413.2820
Local FineWeb-Edu extended train4.134991514.2914
Local FineWeb-Edu train4.166892514.6913

如果你想看完整的数字,它们就在下面。

最初让我惊讶、并决定进行多次LLM判断器运行的数字是”JAX, FineWeb-Edu“模型的那个。在我的第一次运行中,它得到24.35,而OpenAI small权重为24.93——如此接近,以至于我想知道重新运行是否可能超过它们。然而,在额外的四次运行中,它的分数始终低于OpenAI模型,并且差距在某些运行中有所扩大。

那么,FineWeb-Edu是这里的明显赢家吗?也许吧。如果你看一下受污染/无污染对,会发现一些有趣的事情。对于50:50混合,受污染数据集训练的模型得到19.30,而无污染数据集训练的模型得到17.69——相差1.61。对于”策划“数据集,情况更有趣:无污染的得到16.63,而受污染的得到13.58,差3.05分。

记住,污染问题涉及的是模型在训练期间是否看到了保留的测试集。这对于基于该测试集的测试损失来说是个问题,但就IFT测试而言完全无关。

从IFT的角度来看,每对中的受污染和无污染模型看到的训练数据在质量上(至少理论上)基本相同。实际上,无污染运行看到的数据与受污染运行的几乎相同,顺序也相同,只是有些条目被省略了,然后在末尾添加了额外的条目。

这组实验的目的是看数据质量如何影响IFT测试集上的结果。但在策划模型的情况下,一些应该与数据质量无关的事情改变了结果3.05分!

如果仅仅改变模型训练所用同质量数据就能如此剧烈地影响IFT分数,那么要确定数据质量是否真的产生了我们所寻找的效果就变得有点困难了。

另一方面,FineWeb-Edu模型得到24.56,比最接近的其他模型得到的20.50高出4.06分——比我们在两个策划数据集模型之间看到的3.05分差异还要大。而且值得注意的是,得到20.50的模型是”JAX, no MHA bias, no dropout“,它的架构有细微差别——在多头注意力块的输出投影上没有偏置。一个更好的比较可能是”JAX, with MHA bias, no dropout“,它得到17.90分,差异达到惊人的6.66分。

我认为如果不使用不同种子创建的大量不同数据集和不同混合进行非常多次训练运行,就很难确切地搞清楚什么是噪声,什么不是。

然而,那会花费大量时间。我认为在这里最好的做法是将其视为一个相当不错的迹象,表明FineWeb-Edu对IFT评估有帮助,但远非确定。但当然值得指出的是,无论噪声是什么,它的范围至少有3.05分——而FineWeb-Edu模型仅比GPT-2 small差0.63分!所以那里很可能确实有些东西。当然,我们不知道那个模型是否(碰巧)得到了FineWeb-Edu token的最佳平衡,并且永远无法获胜——还是它得到了糟糕的平衡,如果换成更好的平衡实际上会击败GPT-2。所以这当然值得记住。

顺便说一句,策划数据集的结果真的让我惊讶。我曾期望它会是最好的,仅仅因为它几乎肯定包含更多事实。我看了看它对问题的回答——脑海中出现的一种可能性是,它可能对诸如”氯的化学符号是什么“或”谁写了《傲慢与偏见》“这样的问题得到更好的回答,但在不太基于知识的任务上失败。但它在基于事实的问题上也很糟糕:

> 说出《傲慢与偏见》的作者。
> 《傲慢与偏见》的作者是Priscilla Finch。
>
> 氯的周期符号是什么?
> 氯的周期符号是H。

据我所知,许多现实世界的训练运行确实包含(通常是过采样的)高教育性训练数据,就像这个模型的数据集一样。但也许我训练的模型太小了,无法利用它们以这种方式获得的数据——也许这样做并期望好结果,就像要求六岁儿童在学到足够知识去利用它之前记忆东西一样[^1]。值得注意的是,GPT-2 small模型也在那些事实性问题上失败了。

好吧,无论如何:我认为我们在这里有一些有用的结果,所以让我们弄清楚这对下一步意味着什么。

结论

我们在这些实验中得到的结果指向两个有趣的方向。

- 训练集中FineWeb数量与(基于FineWeb的)测试损失评估结果之间的完美联系,虽然事后看来完全明显,但确实凸显了OpenAI small权重在该测试上表现如此之好是多么神秘。
- FineWeb-Edu在IFT测试上表现良好告诉我们,使用更丰富的训练数据似乎确实有价值——尽管50:50混合和策划数据集的表现不那么出色,以及数据选择噪声的指示(来自同等高质量数据集)削弱了这一点。OpenWebText的结果我想我会忽略,因为——虽然理论上它应该与OpenAI训练所用的相似——但没有保证,而且它可能以不显而易见的方式因不显而易见的原因而不同。

我认为向前推进的正确方向是将这两个角度分开。

我应该追求更高的IFT分数,一旦我搞定了这一点,我应该看看什么(如果有的话)可能让由此产生的模型提高其测试分数。但我需要确保无论使用什么数据集,我都要使用多种”混合“——用不同随机种子创建的版本。

在我早期的过度训练实验中,我确实发现它似乎不能改善IFT结果——但它确实改善了测试损失。所以也许找到其他因素的正确组合来提升IFT分数,然后对结果进行过度训练,可能会有帮助?

当然,我的过度训练测试用的是FineWeb,所以如果起始模型(很可能)是在不同数据集上训练的,这种联系可能不太成立。

此外,在处理这里的结果时,我得出的结论是,我使用的模型集合有点令人困惑——现在有不同的超参数设置、小的架构差异(MHA偏置问题)、预训练期间的dropout设置,以及现在的数据集。我认为暂时没问题;我应该把本系列的这部分看作更多的是头脑风暴,而不是实际运行正式实验。但最后,当我有了合理备份的可靠假设时,我应该从头开始:一个基线模型,然后分阶段干预,逐步构建到(希望)一个能和GPT-2 small一样好的模型。

不管怎样,我在这里就此打住。我认为下一个要拉的杠杆(也许令人惊讶)将是权重绑定。我以前有点不把它当作一种可能性,但在我写这个帖子的时候,有东西在我脑海中浮现。OpenAI模型最初是用权重绑定训练的。我的代码库实际上确实支持这样做——但因为我是从”Build a Large Language Model (from Scratch)“的代码中获得的OpenAI权重,所以当我运行IFT测试时,权重实际上并没有绑定!我们加载一个具有独立但相同的嵌入矩阵和输出头矩阵的模型,然后对其进行微调。所以那两个矩阵在微调期间可以独立变化——换句话说,虽然GPT-2 small预训练时有1.24亿参数,但IFT测试是在一个1.63亿参数版本上进行的。这是否给了它们某种不明显的优势?在我的模型上添加权重绑定是否有帮助,无论微调时输出头是否独立?

敬请期待:-)

附录:所有IFT判断器运行

以下是所有IFT判断器运行的数字,为完整起见而包括。你可以看到LLM判断器在运行之间对模型的排名非常一致,但存在变化——也就是说,在某些运行中它处于我认为的”更好的心情“,如果是这样,它会给出更好的分数——但它几乎在模型之间一致地给出,所以所有模型都做得更好。注意(与上面的表不同)这个表按平均IFT分数排序,而不是按测试损失。

模型Run 1Run 2Run 3Run 4Run 5Average
OpenAI weights: medium42.2442.1642.9541.8342.6142.36
OpenAI weights: small24.9324.9625.3925.0125.6625.19
JAX, FineWeb-Edu24.3524.5524.324.6824.924.56
JAX, no MHA bias, no dropout20.519.920.7621.2520.0720.50
JAX, FineWeb/FineWeb-Edu 50:50 (contaminated)19.1618.8619.6119.1719.719.30
JAX, overtrained two normal epochs18.4718.2919.1718.6918.9118.71
JAX, overtrained one long epoch18.0418.7119.6218.4118.5718.67
JAX, with MHA bias, no dropout17.4917.3518.3317.7318.6217.90
JAX, FineWeb/FineWeb-Edu 50:50 (uncontaminated)17.3717.7317.5318.0117.8317.69
JAX, curated (uncontaminated)16.7716.0317.316.0816.9616.63
Cloud FineWeb, 8x A100 40 GiB16.4416.2317.1416.6216.5416.59
1xrtx3090-baseline14.8515.0715.1915.1415.5115.15
Local FineWeb-Edu train14.3714.2315.0814.791514.69
Local FineWeb-Edu extended train14.414.0713.8214.5614.6114.29
Cloud FineWeb, 8x B200 160 GiB13.3713.0513.8513.6714.5713.70
8xa100m40-baseline13.6413.3613.913.3213.9713.64
Cloud FineWeb, 8x H100 80 GiB13.4513.3213.613.5114.0713.59
JAX, curated (contaminated)13.0913.4813.9513.2314.1513.58
1xrtx3090-stacked-interventions13.3713.1114.0413.8413.1713.51
JAX, openwebtext12.8812.713.7413.5313.5313.28
JAX, no MHA bias, with dropout13.1912.8612.9812.8513.2413.02
Local FineWeb train11.7511.7512.2111.4612.1911.87
Cloud FineWeb, 8x A100 80 GiB10.6810.211.0310.5511.4910.79
8xa100m40-stacked-interventions-19.449.7910.8410.210.6610.19

---

[^1]: 一个小男孩右侧卧睡着,右臂伸出,右手无力地耷拉在床沿。透过盒子侧面的圆形格栅,一个声音轻声说话。”尼罗河是非洲最长的河流,也是地球上所有河流中第二长的。虽然不及密西西比-密苏里河的长度,但尼罗河在其流域长度方面居于所有河流之首,其流域延伸跨越35个纬度……“第二天早餐时,有人说:”汤米,你知道非洲最长的河流是哪条吗?“摇摇头。”但你不记得有什么东西开头是:尼罗河是……"“尼罗河——是——非洲——最长的——河流——也是——所有——河流——中——第二——长的……”话脱口而出。“虽然——不及……”“那么,非洲最长的河流是哪条?”眼睛茫然。“我不知道。”“但是尼罗河,汤米。”“尼罗河——是——非洲——最长的——河流——也是——第二……”“那么哪条河最长,汤米?”汤米放声大哭。“我不知道,”他嚎啕道。——《美丽新世界》,奥尔德斯·赫胥黎

阅读原文
📚 相关主题 大语言模型模型训练

订阅 AI Pulse

每天 08:00 · 12:30 · 18:30 · 23:50 更新