Full Width [alt+shift+f] Shortcuts [alt+shift+k]
Sign Up [alt+shift+s] Log In [alt+shift+l]
1
I'm still working on improving the test loss for a from-scratch GPT-2 small base model, trained on code based on Sebastian Raschka's book "Build a Large Language Model (from Scratch)". In my training code, I have this code to create the optimiser: optimizer = torch.optim.AdamW( model.parameters(), lr=0.0004, weight_decay=0.1 ) The values in there -- 0.0004 for the learning rate, and 0.1 for the weight decay -- were just copied from the tiny training run that we do in section 5.2 of the book. What do those values actually mean, and are those really the right values for them? I felt I had a good handle on the learning rate, at least -- it's one of the first things you learn when you start looking at machine learning of any kind -- but how would you go about working out what the correct value for it was? On top of that, when I was reading the Chinchilla paper a while back, I noticed they repeatedly referred to a "cosine cycle" for the learning rate, which didn't fit into anything I'd learned about before. The weight decay was pretty much an unknown for me -- I know it is a parameter controlling the behaviour of the optimiser, but I don't know how it does that. In this post I want to look into the learning rate, and these mysterious cosines; I'll write a follow-up about the weight decay later. The learning rate: a refresher If you're reading this blog, you almost certainly know what the learning rate is, but let's go over it briefly to build a solid foundation. The way it's normally explained, using simple gradient descent, goes something like this. Let's assume that we're training a model with just one parameter, and it starts off set to −5. We run some training data through, and get a loss, let's say 44.44: We don't know what shape our loss curve is (if we did, we might be able to find the lowest loss algebraically), but we do know the differential of the parameter versus the loss at the point we've measured; it happens to be -13. That is...
10th Mar 2026

Stay updated

Get a weekly newsletter with the top 5 articles worth reading every week.

More from Giles' blog

Why do OpenAI's GPT-2 weights beat mine? Part five: data quality

When I finished learning how to build an LLM from scratch, I was left with a mystery: my own models were not as good as OpenAI's original GPT-2 models, despite being based on the same architecture. My models all had 163M parameters, and followed the design from Sebastian Raschka's book "Build a Large Language Model (from Scratch)". That meant that they were pretty much the same as the setup for the OpenAI GPT-2 "small" instance, except that they did not use weight-tying or bias on the QKV matrices. Weight-tying means that you re-use the initial embedding matrix as the output head at the end, and using it means that GPT-2 small saved quite a few parameters -- it was 124M rather than 163M -- at, at least in my own experiments, a cost in quality; similarly, while I found that QKV bias made a tiny improvement in loss terms, I'd felt it was likely within the noise. But GPT-2 small consistently beat my models on an instruction fine-tuning (IFT) task -- also adapted from Raschka's book. That test fine-tunes the model on a subset of the Alpaca dataset, until validation loss starts rising, and then runs a test set through the resulting model. The responses to the test set questions are stored, and then I run all of the responses from all of the models under test past GPT 5.5 in one go to get an aggregate score; more details here. GPT-2 small always did better than any of my models on this. Additionally, it did surprisingly well on a simpler eval -- one that just measured the cross entropy loss it got on a test set. It scored close to my own best models, and better than many of them. What made this result particularly interesting was that the test set in question was a split of my own training data; my models would not have seen it when training (at least, in theory), but it seems likely that it would be much more similar to their own training data than it was to OpenAI's. I've checked two things while probing this mystery: It seems very likely that the GPT-2 models were overtrained by modern standards; would overtraining my own models get them closer? It turned out that no, it probably didn't help with the IFT eval (though there might have been some signal there). It did help quite a lot with the test loss eval, though. The way I was handling dropout in the IFT test might have been unduly benefiting some models while working against others. I decided to standardise on not using dropout during this eval, as (counter-intuitively for me) it seemed to harm the results of most models, even those that had been pre-trained with dropout. In particular, the OpenAI weights were harmed by using dropout, and making a change that benefited them (along with some of my own models) seemed the most conservative approach to take in investigating this. The next thing I wanted to look into was the training data. The exact dataset that the various GPT-2 models were trained on has never been released; all we know about it is from the paper, where they say: [W]e created a new web scrape which emphasizes document quality. To do this we only scraped web pages which have been curated/filtered by humans. Manually filtering a full web scrape would be exceptionally expensive so as a starting point, we scraped all outbound links from Reddit, a social media platform, which received at least 3 karma. This can be thought of as a heuristic indicator for whether other users found the link interesting, educational, or just funny. They called it "WebText". There is an OpenWebText that tries to replicate it, but although they tried to follow the same procedure as the original, there's no guarantee that it is all that similar. By comparison, I'd normally been training against FineWeb. While this is a general web-scraping dataset, without the "curation" provided by using only stuff that was linked from upvoted Reddit posts, it has been refined to remove any obvious junk. I had felt that it was pretty much equivalent. But what if I were wrong about that? I decided to see if I could get better models by using better data. The starting point Here's a table of all of the models I've been comparing to date. The "Test loss" column shows how well the model in question did on that held-back cross entropy loss evaluation. The "IFT epochs" column shows how many epochs of fine-tuning the model needed before its validation loss started rising, the "IFT score" the score that GPT 5.5 gave the model's responses to the test set of my Alpaca data, and the "IFT rank" the model's rank in terms of that score. The OpenAI small model is in there in bold, and I've also included the OpenAI medium model for comparison purposes. Test loss IFT epochs IFT score IFT rank OpenAI weights: medium 3.231442 2 43.75 1 JAX, overtrained one long epoch 3.324953 3 19.77 4 JAX, overtrained two normal epochs 3.326482 4 19.72 5 JAX, with MHA bias, no dropout 3.418784 4 18.69 6 JAX, no MHA bias, no dropout 3.420089 5 21.46 3 JAX, no MHA bias, with dropout 3.476802 5 13.22 15 OpenAI weights: small 3.499677 2 26.00 2 1xrtx3090-stacked-interventions 3.538161 4 13.77 14 8xa100m40-stacked-interventions-1 3.577761 4 10.76 18 Cloud FineWeb, 8x A100 40 GiB 3.673623 3 17.72 7 1xrtx3090-baseline 3.683835 4 15.74 8 8xa100m40-baseline 3.691526 3 14.19 13 Cloud FineWeb, 8x H100 80 GiB 3.724507 4 14.33 12 Cloud FineWeb, 8x A100 80 GiB 3.729900 3 11.34 17 Cloud FineWeb, 8x B200 160 GiB 3.771478 4 14.67 11 Local FineWeb train 3.943522 5 12.31 16 Local FineWeb-Edu extended train 4.134991 5 15.04 9 Local FineWeb-Edu train 4.166892 5 14.99 10 You can see that the OpenAI small model did pretty well in terms of the test loss, when you consider that it has 39M fewer weights than my models and was being tested against a dataset that differs more from its likely training data than it does from my own models'. Additionally, the specific models that did better than OpenAI's small one were all trained with JAX rather than PyTorch -- my hypothesis for that is that it's a result of the JAX ones getting better initial weights by pure chance. But the big difference was in the IFT score. In the specific run that gave the results in this table, the OpenAI small model got 26.00 -- the closest of my own models was more than 4.5 points lower, at 21.46. This difference was consistent over all of my other test runs. The GPT-2 small model was always ahead of mine. (GPT-2 medium, of course, beat GPT-2 small and all of my models, but given that it is twice the size of mine, that's not a big surprise.) Now, quite some time ago, I had tried looking into data quality as a lever to pull for model performance. At the bottom of the table, with the worst test loss of all models, you can see two models: "Local FineWeb-Edu train" "Local FineWeb-Edu extended train" These two were (as you might guess from the names) trained on the FineWeb-Edu dataset, which includes just the most "educational" data from FineWeb. They scored very badly on the test loss score. Given that the test dataset is from FineWeb, that's not a big surprise -- as I've written previously: If you train a model on Jane Austen and then evaluate against Chuck Tingle, then you're not going to get amazing results. But again, GPT-2 had the same issue, and did perfectly well on the test loss eval. On the other hand, while these FineWeb-Edu models' performance on the IFT eval wasn't stellar -- there are plenty of my other models ahead of them -- they did seem to punch above their weight. Consistently across all of the IFT evals I've done, they have scored higher than many of the others -- despite their poor loss on the test eval. Additionally: they were amongst the first models that I trained, before I'd spent time learning about how to optimise my hyperparameters and training loop. They did not use gradient clipping, they did use dropout, their batch size was just "whatever I could squeeze into the GPU", and I didn't set the learning rate to the right kind of value or schedule it over the course of the training run. So maybe a new training run on FineWeb-Edu plus my training improvements would help? And maybe some other tweaks to the training data would be worth looking into? The plan I decided to see what would happen if I trained some models with better-quality data. Specifically, I would train models with my current optimised loop and hyperparameters on four different datasets: FineWeb-Edu -- essentially the same as "Local FineWeb-Edu train" but with a better training setup. This would test the "more educational -> better" hypothesis. A 50:50 split of FineWeb and FineWeb-Edu. I've read that LLMs can be helped by having a decent amount of lower-quality data in their training loop, as it helps them to generalise. Perhaps having some FineWeb in there in addition to the FineWeb-Edu stuff would improve that test loss score while also helping the IFT test? A "curated" dataset containing 45% of its contents from FineWeb, 45% from FineWeb-Edu, and 10% from the Simple English Wikipedia. The full Wikipedia is huge, and full of obscure facts -- while the Simple English one is small and hopefully richer in useful information on a per-token basis. And conveniently, Answer.ai have made a snapshot of it available on Hugging Face Hub. Might deliberately putting a bunch of encyclopaedic data into the training set make the model better at the IFT eval (which has lots of factual questions in it, like "who wrote Pride and Prejudice")? OpenWebText. Even though I was unsure how well it matched the original WebText, given that it was there, it seemed silly to not try training something on it and see how it matched up. I would train each model on 3.2B tokens of the chosen dataset; that's the Chinchilla-optimal amount for my 163M-parameter models. If there were any interesting results, then I might consider doing overtrained models later on. I decided to be at least vaguely scientific about this, and to pre-register some predictions: The FineWeb-Edu-only model would do pretty badly on the test loss, but better than my older FineWeb-Edu models (90%). It would also punch above its weight on the IFT eval (90%). The 50:50 split: I expected it to do worse on the test eval than my JAX FineWeb-only models (70%), but better than the FineWeb-Edu one (90%). I wasn't sure about how it would do on the IFT eval, but thought it might be somewhere in between the two groups (60%). The curated dataset I had high hopes for in terms of the IFT eval -- let's say 80% chance of it being the best of all of my models. For the test loss eval, I expected it to do about as well as the 50:50 split, maybe a little bit worse (70%). I had no idea how the OpenWebText eval would do! Could be worse, could be better. Here's how things turned out. The FineWeb-Edu model I already had a dataset based on FineWeb-Edu ready to go, from when I trained those two original models. It is just the 10B-token sample of the original dataset at the time I generated it last December, formatted appropriately for my training script (details on the dataset card). I kicked off a training run with my JAX code (which I've been using for the other posts in this series): 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] ...and just less than 40 hours later, I had a model: 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 I converted the saved JAX safetensors file from the last checkpoint into a format that would be compatible with my PyTorch eval code, and ran my smoke test: how would it complete the sentence "Every effort moves you"? Every effort moves you closer to God’s Kingdom, and even closer to Him. As we can see in That was nice and coherent -- if unusually religious! -- so that was promising. I ran the test eval: 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 That was pretty good, putting it at a better test loss than all of the models I had trained without optimised hyperparameters, and worse than all of the ones I had trained on FineWeb with optimised hyperparameters. So that fit in with my prediction that it would be better than the old FineWeb-Edu models; the fact that it was also better than the non-optimised training runs with FineWeb seemed sensible enough that I felt silly for not having predicted that it would have fallen exactly there :-) I decided to leave the IFT eval until the end so that I could check all of the models from these experiments together, so it was time to upload this one to Hugging Face, and move on to the next model. 50:50 FineWeb to FineWeb-Edu I put together a new repo with a script to prepare datasets specifically for my training setup. You provide it with config that specifies some source datasets along with information about how to process them and how to mix them together, and it uploads a new dataset to Hugging Face Hub with the required characteristics. For example, for the 50:50 FineWeb to FineWeb-Edu split, the config looked like this: { "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 } ] } The way the script works is pretty simple: it works out (based on those weights and the tokens_desired) how many tokens it wants from each source dataset, shuffles the items in the sources, then it loops until it has the desired number of tokens or more stored in an output. In the loop, it works out which source is currently most under-represented, grabs an item from it, tokenises it, and adds it to the output. Running it with that 50:50 config seemed to work fine: 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 So we had almost-perfect 50:50 balance between the datasets, and it saved this dataset on Hugging Face. I ran a script to double-check that it looked sane, and it did, so it was time to spin up a training run: 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] That was running on perry, my normal workstation, and I kicked it off in parallel with the "curated" model training run below on poppy my training box, but I'll keep the runs separate for the purposes of this writeup. When this had been running for an hour or so, our power went out. My guess is that having the tumble dryer running, the car charging, the kettle boiling, the electric hob switched on, and two machines doing training runs is a bit too much for our electrics... which might be a problem in the future, especially if (as planned) I make poppy a multi-GPU machine. However, as things stand, I was able to kick it off again after switching the circuit breaker back on, and things held up. Again, about 40 hours later: 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 (Note that the numbers reported at the end of a restarted run like this only include what happened after the restart.) I converted it to PyTorch-compatible tensors, and did the smoke test: Every effort moves you on to other options—in fact, it’s not even worth that effort. Just make Looking good! Time for the loss test: 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 That was almost in keeping with my prediction that it would do worse than the JAX FineWeb-only models, except that it was better than the worst of those, "JAX, no MHA bias, with dropout": it was actually better than I predicted. So, a promising model. Time to upload it to Hugging Face -- and now let's move on to the next one. The "curated" dataset With my dataset-preparation script, this was easy enough to set up: { "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 } ] } Running that worked nicely: 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 One thing that is worth noting in that output is the "6 iterators" for the Simple English Wikipedia. If a source dataset runs out of items while we're building up the results in this script, we start iterating over it again (with a different seed for the shuffle so that the ordering is different). The "6 iterators" means that it needed to do that 6 times -- the original creation of the iterator at the start of the script, and five more. So that means that the Simple English Wikipedia is repeated (oversampled) somewhere between five and six times in the dataset. That's not a bad thing! From what I've read, it's actually quite standard to oversample highly educational content in LLM training datasets. And anyway, the dataset the script generated was 10B tokens, of which we're only using 3.2B for the training run in this post, so it would only appear somewhere between one and two times. The repetition would likely only really cut in if and when we did an overtrained model on the dataset. Anyway, I ran my check against the uploaded dataset -- the first few items were clearly from FineWeb, FineWeb-Edu, and the Simple English Wikipedia. It was time to kick off a training run: 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] Again, this was interrupted by the power outage that hit the 50:50 training run, but I was able to restart from a checkpoint. After another 22 hours, it crashed with an error that I've seen before: 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'] I put it aside as a one-off oddity when I hit it last time, but this time I dug in a bit more. I noted that it had not ever happened on perry, but seemed to be an issue on poppy, and that poppy had an older version of CUDA and the Nvidia drivers -- might that be the cause? I decided to upgrade those before kicking off the next run, but for now just restarted the run from the most recent checkpoint. (Note for anyone who is hitting the same error: it has not occurred since the upgrade, so that's worth trying.) This time it completed OK: 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 Again, these numbers just show what happened after the most recent restart. I copied it over to perry, converted it into a format that was compatible with my PyTorch code, and ran the smoke test: Every effort moves you by the air, for it will make you a better athlete, so your body becomes bigger and stronger Coherent enough -- time for the loss eval: 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 Again, in line with my predictions -- worse than the JAX FineWeb-only models, and indeed than the very best PyTorch one, 1xrtx3090-stacked-interventions, and also worse than the 50:50 split, but better than the FineWeb-Edu one. I uploaded it to Hugging Face, and it was time to move on to what was meant to be the final model for this set of experiments. The OpenWebText run Again, this was a simple enough config to set up: { "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 } ] } ...and the build and upload process worked well (and took much less time -- for some reason, sampling randomly from a single dataset is faster than sampling from two or three): 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 Note that it needed to oversample -- that "2 iterators". OpenWebText is about 40 GiB uncompressed, and so that's about 10B GPT-2 tokens -- presumably just a little bit less. Again, given that I was planning to use just the first 3.2B tokens of the dataset, I didn't feel that it would matter. I ran the check script on the newly-uploaded Hugging Face dataset and all looked well, so that was all set for the training run. I upgraded poppy first with a sudo pacman -Syu to see if that helped with the weird error that I got in the previous run (which, as I said, it looks like it did), then kicked it off: 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] About 31 hours in, it crashed again, but this time it was my own dumb fault: poppy has a relatively small disk and I ran out of space. I fixed that and kicked it off again from the most recent checkpoint, and this time it completed: 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 I converted it to PyTorch for the smoke test: Every effort moves you through each phase, so it's not a complete picture. I'm sure your story was ...which looked solid, so it was time for the test loss eval: 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 Our worst score yet in this experiment! Worse than any of my models so far, apart from the two FineWeb-Edu ones I did without optimised hyperparameters. Now, the first draft of this post went straight to the results from here, but the story wasn't quite over yet... Test set contamination GPT-6 Astra is relentless. Before I publish any of these posts, I run them past an editorial board of LLMs to look for issues. GPT-6 Astra not only checked the text, it also visited the code I'd linked to to check that out too, and spotted something problematic. It's obvious in retrospect, but my code to build the new datasets had a high risk of including the contents of the -- in theory held-back -- test set. The way that the test set was generated was that I downloaded the 10B sample of FineWeb back in December, splitting it into 99% training data and 1% "validation". That validation split was about 100M tokens, and I was only using the first 19M or so for actual validation runs during training, so I (somewhat arbitrarily) designated about 19M other tokens starting at position 50M in there as my test set. Now, my new dataset-generation code was just sampling randomly from the complete 10B sample of FineWeb. So there was nothing stopping it from pulling in data that was in that old validation split! That meant that it was quite likely that my new "curated" and "50:50" datasets contained at least some of the test set that was meant to have been held back from the models during training. On reflection, the problem was potentially even worse. FineWeb-Edu is a subset of FineWeb; my existing FineWeb-Edu dataset came from the 10B sample of the Hugging Face original, and so it also could potentially contain documents that I'd put into the test set. The first thing to do was to establish the size of the problem. I wrote a script to take in a "forbidden" dataset and split; this was assumed to be formatted as one big tensor of GPT-2 tokens, which is what all of my datasets are. It would then split it by end-of-text tokens, and generate a hash and a token count for each resulting "document". Optionally, you could restrict it to only considering a subset -- the n tokens starting at position p -- and it would then generate hashes/lengths for the documents inside that slice, or that overlapped it at the start or the end. I ran that to generate a list of hashes for the entire validation set -- the validation split of gpjt/fineweb-gpt2-tokens -- and then used a second script to check my various training sets (and the validation set itself) to see how much of a contamination problem there was. I got these results: Dataset Split Contamination with validation set gpjt/fineweb-gpt2-tokens validation 102163003 out of 102163003 tokens (100.00%) gpjt/fineweb-gpt2-tokens train 636166 out of 102163003 tokens (0.62%) gpjt/fineweb-edu-gpt2-tokens train 672189 out of 102163003 tokens (0.66%) gpjt/fw-fwedu-5050-gpt2-tokens train 49224580 out of 102163003 tokens (48.18%) gpjt/fw-fwedu-simplewiki-gpt2-tokens train 44233824 out of 102163003 tokens (43.30%) gpjt/openwebtext-gpt2-tokens train 212 out of 102163003 tokens (0.00%) So: The validation set was 100% "contaminated" with itself, which was a useful sanity check. The training set of gpjt/fineweb-gpt2-tokens had what I felt was a small level of contamination. It was interesting that there was any at all -- I think that must mean that there are some repeated documents in the original dataset, and some of them wound up with copies in both my training and validation splits. The gpjt/fineweb-edu-gpt2-tokens dataset also had what felt like a reassuringly low level of contamination. Both gpjt/fw-fwedu-5050-gpt2-tokens and gpjt/fw-fwedu-simplewiki-gpt2-tokens, however, looked problematic. In both cases, the training datasets had more than 40% of the validation/test set in them. gpjt/openwebtext-gpt2-tokens was, as you'd expect, almost completely uncontaminated. It looks like maybe one document happened to have been picked up by both the OpenWebText and the FineWeb crawls and then included in the bit of FineWeb I was using for validation. However, these numbers -- while scary, at least for the 50:50 and the curated datasets -- were not quite the ones to use. They showed how much of the full validation set showed up in the full training set; what I actually cared about was how much of the test set -- those 19M tokens starting at position 50M in the validation split -- was in the actual subset of the training datasets that I actually trained on -- the first ~3.2B of them. I re-ran the script to generate hashes for just the test set, and then re-ran the contamination-checking script, telling it just to look at the appropriate subset of the training tokens, and got this: Dataset (first 3.2B tokens only) Split Contamination with test set gpjt/fineweb-gpt2-tokens train 26557 out of 19632681 tokens (0.14%) gpjt/fineweb-edu-gpt2-tokens train 32079 out of 19632681 tokens (0.16%) gpjt/fw-fwedu-5050-gpt2-tokens train 2986889 out of 19632681 tokens (15.21%) gpjt/fw-fwedu-simplewiki-gpt2-tokens train 2682430 out of 19632681 tokens (13.66%) gpjt/openwebtext-gpt2-tokens train None It was clear that there was a problem -- certainly with gpjt/fw-fwedu-5050-gpt2-tokens and gpjt/fw-fwedu-simplewiki-gpt2-tokens. They'd seen what felt like a significant amount of the test set while training, so their results on the test loss eval were dubious at best. I decided to train those two models afresh, and see what the result was in terms of loss. If the difference was huge, I'd look into the risks of the (much smaller) contamination of gpjt/fineweb-gpt2-tokens and gpjt/fineweb-edu-gpt2-tokens. But if it was pretty small, I'd not worry about that too much. I extended the script that prepared datasets so that the config file could specify a forbidden_dataset. Any documents in the source datasets that matched forbidden ones would be excluded from the output. I then updated the config for gpjt/fw-fwedu-5050-gpt2-tokens and gpjt/fw-fwedu-simplewiki-gpt2-tokens so that the whole validation split of gpjt/fineweb-gpt2-tokens was forbidden, and re-generated them. You can see the updated datasets here and here. Running the contamination-checker script against them showed that they were clear. I then re-did the full training runs for those models; the uncontaminated version of the 50:50 split model is here, and the curated one is here. And the good news: both of them actually did very slightly better at the test loss eval than their equivalents that had been trained on the contaminated data: Model Contaminated Test loss JAX, FineWeb/FineWeb-Edu 50:50 No 3.449257 JAX, FineWeb/FineWeb-Edu 50:50 Yes 3.462454 JAX, curated No 3.534068 JAX, curated Yes 3.542460 There are a number of possibilities that come to mind; perhaps learning from the test set just doesn't happen with tiny 163M models like this, or perhaps while the contaminated models were learning, the benefit they got from that was outweighed by the data that they got instead of the test set data being in some way better for training purposes, at least in terms of the loss eval. But anyway, I felt that if the effect of seeing more than 10% of the test set data during training was so tiny, then the effect of seeing less than 0.2% -- which is what the FineWeb-Edu model in this set of training runs had, as did all of my other FineWeb-only models from previous experiments -- would be even smaller and I'd disregard it. That was excellent news! I didn't need to start all of my experiments from scratch. For the rest of this post, I will include the numbers and results for the contaminated models as well as the uncontaminated ones -- they're interesting for several reasons -- but for future posts I'll skip the contaminated ones. So -- finally! -- let's start digging into the final results. Results Firstly, I think it's worth taking a look at all of the test loss results in context. Here they are in a table, with the new models in bold: Test loss OpenAI weights: medium 3.231442 JAX, overtrained one long epoch 3.324953 JAX, overtrained two normal epochs 3.326482 JAX, with MHA bias, no dropout 3.418784 JAX, no MHA bias, no dropout 3.420089 JAX, FineWeb/FineWeb-Edu 50:50 (uncontaminated) 3.449257 JAX, FineWeb/FineWeb-Edu 50:50 (contaminated) 3.462454 JAX, no MHA bias, with dropout 3.476802 OpenAI weights: small 3.499677 JAX, curated (uncontaminated) 3.534068 1xrtx3090-stacked-interventions 3.538161 JAX, curated (contaminated) 3.542460 8xa100m40-stacked-interventions-1 3.577761 JAX, FineWeb-Edu 3.632900 Cloud FineWeb, 8x A100 40 GiB 3.673623 1xrtx3090-baseline 3.683835 8xa100m40-baseline 3.691526 Cloud FineWeb, 8x H100 80 GiB 3.724507 Cloud FineWeb, 8x A100 80 GiB 3.729900 Cloud FineWeb, 8x B200 160 GiB 3.771478 Local FineWeb train 3.943522 JAX, openwebtext 4.045255 Local FineWeb-Edu extended train 4.134991 Local FineWeb-Edu train 4.166892 I think there's something very clear here: with the new models, the more FineWeb that was in the training mix, the better the model did on this eval. I think I might have been subconsciously expecting that in the predictions I did before running these experiments, but in retrospect it's so incredibly obvious that I feel silly for not mentioning it explicitly! But that tells us something interesting. From the description in the paper, whatever OpenAI did the GPT-2 training run on, it was not like FineWeb. It was probably more similar to OpenWebText -- and yet, that model was the one that performed the worst on this test eval, so if it is more like OpenWebText, there must be some other factor involved. But moving on for now: how about the IFT test -- the one that kicked off all of this work in the first place? I generated a set of IFT responses for all of the new models, and then ran them (plus responses for all of the other models on that table above) past GPT 5.5, and found that one of my new models was getting quite close to the original GPT-2 small weights! So I did four more runs, so that I could get an average. Here are the results -- the "IFT score" is the average across all five runs of the judge, and the "IFT rank" is based on that. The "IFT epochs" was from the original result-generation script. Test loss IFT epochs IFT score IFT rank OpenAI weights: medium 3.231442 2 42.36 1 JAX, overtrained one long epoch 3.324953 3 18.67 7 JAX, overtrained two normal epochs 3.326482 4 18.71 6 JAX, with MHA bias, no dropout 3.418784 4 17.90 8 JAX, no MHA bias, no dropout 3.420089 5 20.50 4 JAX, FineWeb/FineWeb-Edu 50:50 (uncontaminated) 3.449257 4 17.69 9 JAX, FineWeb/FineWeb-Edu 50:50 (contaminated) 3.462454 4 19.30 5 JAX, no MHA bias, with dropout 3.476802 5 13.02 21 OpenAI weights: small 3.499677 2 25.19 2 JAX, curated (uncontaminated) 3.534068 4 16.63 10 1xrtx3090-stacked-interventions 3.538161 4 13.51 19 JAX, curated (contaminated) 3.542460 4 13.58 18 8xa100m40-stacked-interventions-1 3.577761 4 10.19 24 JAX, FineWeb-Edu 3.632900 4 24.56 3 Cloud FineWeb, 8x A100 40 GiB 3.673623 3 16.59 11 1xrtx3090-baseline 3.683835 4 15.15 12 8xa100m40-baseline 3.691526 3 13.64 16 Cloud FineWeb, 8x H100 80 GiB 3.724507 4 13.59 17 Cloud FineWeb, 8x A100 80 GiB 3.729900 3 10.79 23 Cloud FineWeb, 8x B200 160 GiB 3.771478 4 13.70 15 Local FineWeb train 3.943522 5 11.87 22 JAX, openwebtext 4.045255 4 13.28 20 Local FineWeb-Edu extended train 4.134991 5 14.29 14 Local FineWeb-Edu train 4.166892 5 14.69 13 If you want to see the full numbers, they're below. The number that initially surprised me, and made me decide to do multiple LLM-judge runs was the one for the "JAX, FineWeb-Edu" model. In my first run it came in at 24.35 vs the OpenAI small weights' 24.93 -- so close that I wondered if it might even beat them on a re-run. However, in the further four runs its score was consistently lower than the OpenAI model's, and the gap extended a bit in some. So, was FineWeb-Edu the clear winner here? Perhaps. If you look at the contaminated/uncontaminated pairs, something interesting pops out. For the 50:50 mix, the model trained with the contaminated dataset got 19.30, and the one trained on the uncontaminated one got 17.69 -- a difference of 1.61. For the "curated" dataset, the situation was even more interesting: uncontaminated got 16.63, while contaminated got 13.58, a delta of 3.05 points. Remember, the contamination issue is about whether or not the model saw the held-back test set during training. It was an issue for the test loss that is based on that test set, but is entirely orthogonal to the IFT test. From the IFT perspective, both contaminated and uncontaminated models in each case saw training data that was -- in theory, at least -- essentially the same in terms of quality. Indeed, the uncontaminated run saw almost the same data in the same order as the contaminated one, except that some items were omitted, and then extra ones were added to the end. The purpose of this set of experiments was to see how data quality affected the results on the IFT test set. But in the case of the curated model, something that should be unrelated to data quality changed the results by 3.05 points! If something as simple as changing which data of the same quality the model is trained with can affect the IFT score so drastically, it makes it a bit harder to be certain as to whether or not data quality really had the effect we were looking for. On the other hand, the FineWeb-Edu model came in at 24.56, which is 4.06 points better than the 20.50 that the closest other model got -- more than the 3.05 points we see in difference between the two curated dataset models. And it's worth noting that the model with 20.50 is "JAX, no MHA bias, no dropout", which has a subtly different architecture -- no bias on the output projection of the multi-head attention blocks. A better comparison might be "JAX, with MHA bias, no dropout", which got a score of 17.90, for a whacking great difference of 6.66 points. I think that without doing a very large number of training runs on different datasets with different mixes, each one created with a different seed, it would be hard to work out exactly what is in the noise here and what is not. However, that would cost a lot in terms of time. I think that the best thing here is to chalk this up as a fairly decent indication that FineWeb-Edu improves matters for the IFT eval, but far from a certainty. But it's certainly worth noting that whatever the noise is, it has a range of at least 3.05 points -- and the FineWeb-Edu model is just 0.63 points short of GPT-2 small! So there could well be something there. Of course, we don't know whether that model got (by chance) the best possible balance of FineWeb-Edu tokens, and could never win -- or whether it got a bad balance and would actually beat GPT-2 with a better one. So that's certainly worth keeping in mind. As an aside, the result for the curated dataset really surprised me. I had expected that it would be the best one, simply because it almost certainly contained more facts. I took a look at its answers to the questions -- one possibility that came to mind might be that it would get better responses to questions like "What is the chemical symbol for chlorine" or "Who wrote Pride and Prejudice" than the others, but would fail on less knowledge-based tasks. But it was terrible at fact-based questions too: Name the author of 'Pride and Prejudice'. What is the periodic symbol for chlorine? As I understand it, many real-world training runs do include (often oversampled) amounts of highly educational training data like this model's dataset did. But perhaps the models that I'm training are just too small to be able to make use of the data they gained that way -- maybe doing things this way and expecting good results is like asking six-year-old children to memorise stuff before they've learned enough to be able to make use of it 1. It's worth noting that the GPT-2 small model also failed on those factual questions. Well, anyway: I think we have some useful results here, so let's work out what that means for next steps. Conclusion The results we got in these experiments point in two interesting directions. The perfect connection between the amount of FineWeb in the training set and the result on the (FineWeb-based) test loss eval, while perfectly obvious in retrospect, really does highlight how mysterious it is that the OpenAI small weights do so well on that test. The fact that FineWeb-Edu did well on the IFT test tells us that there does seem to be value in using richer training data -- though the less-spectacular results of the 50:50 mix and the curated one weaken that a bit, as does the indicator of what the noise due to data selection from equivalently high-quality datasets might be. The OpenWebText result I think I'll ignore, given that -- while in theory it should be similar to what OpenAI trained on -- there are no guarantees, and it might differ in non-obvious ways for non-obvious reasons. I think that the right direction to take this going forward is to separate these two angles. I should chase a higher IFT score, and then once I have nailed that down, I should see what (if anything) might allow me to get the resulting model to improve its test score. But I will need to make sure that whatever dataset I use, I use various "mixes" of it -- versions created with different random seeds. In my earlier experiments with overtraining, I did find that it didn't seem to improve the IFT results -- but it did improve the test loss. So perhaps identifying the right combination of other factors to boost the IFT score, then overtraining the result, might help? Of course, my overtraining tests were with FineWeb, so the connection might not hold up as well if the starting model (as seems likely) was trained on a different dataset. Also, while working through the results here, I've come to the conclusion that the set of models I'm using is a bit confusing -- there are now different hyperparameter settings, small architectural differences (the MHA bias thing), dropout settings during the pre-training, and now datasets. I think that's OK for now; I should see this part of this series as more ideation than actually running the proper experiments. But at the end, when I have some solid hypotheses with a reasonable amount of backup, I should start from scratch: a baseline model, then staged interventions to build up to what (hopefully) will be a model as good as GPT-2 small. Anyway, I'll wrap this one up here. I think that the next lever to pull is (perhaps surprisingly) going to be weight tying. I had previously kind of disregarded that as a possibility, but while I was working on this post, something popped into my mind. The OpenAI models were originally trained with weight tying. My codebase does actually support doing it -- but because I got the OpenAI weights I'm using from the code in "Build a Large Language Model (from Scratch)", when I'm running the IFT test, the weights are not actually tied! We load up a model that has separate but identical embedding and output head matrices, and then we fine-tune that. So those two matrices can vary independently during fine-tuning -- to put it another way, while GPT-2 small was pre-trained with 124M parameters, the IFT test is being done on a 163M-parameter version. Does that give them some non-obvious advantage? And would adding weight-tying to my own models help, either with or without the output heads being independent at fine-tuning time? Stay tuned :-) Appendix: all IFT judge runs Here are the numbers for all of the IFT judge runs, included for completeness. You can see that the LLM judge ranks models very consistently between runs, but there is variation -- that is, on some runs it's in what I think of as a "better mood" than others, and if that's the case, it will give better scores -- but it will give them almost consistently between models, so all of the models do better. Note that (unlike the table above) this one is sorted by the average IFT score rather than the test loss. Model Run 1 Run 2 Run 3 Run 4 Run 5 Average OpenAI weights: medium 42.24 42.16 42.95 41.83 42.61 42.36 OpenAI weights: small 24.93 24.96 25.39 25.01 25.66 25.19 JAX, FineWeb-Edu 24.35 24.55 24.3 24.68 24.9 24.56 JAX, no MHA bias, no dropout 20.5 19.9 20.76 21.25 20.07 20.50 JAX, FineWeb/FineWeb-Edu 50:50 (contaminated) 19.16 18.86 19.61 19.17 19.7 19.30 JAX, overtrained two normal epochs 18.47 18.29 19.17 18.69 18.91 18.71 JAX, overtrained one long epoch 18.04 18.71 19.62 18.41 18.57 18.67 JAX, with MHA bias, no dropout 17.49 17.35 18.33 17.73 18.62 17.90 JAX, FineWeb/FineWeb-Edu 50:50 (uncontaminated) 17.37 17.73 17.53 18.01 17.83 17.69 JAX, curated (uncontaminated) 16.77 16.03 17.3 16.08 16.96 16.63 Cloud FineWeb, 8x A100 40 GiB 16.44 16.23 17.14 16.62 16.54 16.59 1xrtx3090-baseline 14.85 15.07 15.19 15.14 15.51 15.15 Local FineWeb-Edu train 14.37 14.23 15.08 14.79 15 14.69 Local FineWeb-Edu extended train 14.4 14.07 13.82 14.56 14.61 14.29 Cloud FineWeb, 8x B200 160 GiB 13.37 13.05 13.85 13.67 14.57 13.70 8xa100m40-baseline 13.64 13.36 13.9 13.32 13.97 13.64 Cloud FineWeb, 8x H100 80 GiB 13.45 13.32 13.6 13.51 14.07 13.59 JAX, curated (contaminated) 13.09 13.48 13.95 13.23 14.15 13.58 1xrtx3090-stacked-interventions 13.37 13.11 14.04 13.84 13.17 13.51 JAX, openwebtext 12.88 12.7 13.74 13.53 13.53 13.28 JAX, no MHA bias, with dropout 13.19 12.86 12.98 12.85 13.24 13.02 Local FineWeb train 11.75 11.75 12.21 11.46 12.19 11.87 Cloud FineWeb, 8x A100 80 GiB 10.68 10.2 11.03 10.55 11.49 10.79 8xa100m40-stacked-interventions-1 9.44 9.79 10.84 10.2 10.66 10.19 A small boy asleep on his right side, the right arm stuck out, the right hand hanging limp over the edge of the bed. Through a round grating in the side of a box a voice speaks softly. "The Nile is the longest river in Africa and the second in length of all the rivers of the globe. Although falling short of the length of the Mississippi-Missouri, the Nile is at the head of all rivers as regards the length of its basin, which extends through 35 degrees of latitude …" At breakfast the next morning, "Tommy," some one says, "do you know which is the longest river in Africa?" A shaking of the head. "But don't you remember something that begins: The Nile is the …" "The - Nile - is - the - longest - river - in - Africa - and - the - second - in - length - of - all - the - rivers - of - the - globe …" The words come rushing out. "Although - falling - short - of …" "Well now, which is the longest river in Africa?" The eyes are blank. "I don't know." "But the Nile, Tommy." "The - Nile - is - the - longest - river - in - Africa - and - second …" "Then which river is the longest, Tommy?" Tommy burst into tears. "I don't know," he howls. Brave New World, Aldous Huxley ↩

3 days ago • 1 votes
Putting my JAX-trained models on the Hugging Face Hub

I hadn't uploaded the models that I trained using JAX to the Hugging Face Hub because Transformers has been PyTorch-only since version 5 (though they say they're working to add interoperability with JAX in the future), so it would have been tough to get them working natively with AutoModelForCausalLM and the like. But then it dawned on me that I'd already written a conversion script that could take my JAX safetensors files and convert them into ones compatible with my PyTorch code. It's actually those converted models that I use for my evals -- so I could use my existing PyTorch script to upload them. So, I've now uploaded PyTorch-compatible versions of all of my JAX-trained models: "Writing an LLM from scratch, part 34b -- from bigrams to GPT-2, one component at a time (in JAX)" gpjt/jax-no-mha-bias-no-dropout -- the first full LLM trained in the post, in the "Adding LayerNorm" section. gpjt/jax-no-mha-bias-with-dropout -- the second full LLM trained in the post, in the "Dropout" section. gpjt/jax-with-mha-bias-no-dropout -- the third full LLM trained in the post, in the "Adding bias to the MHA output projections" section. "Why do OpenAI's GPT-2 weights beat mine? Part three: testing overtraining" gpjt/jax-with-mha-bias-no-dropout-extended -- the single-epoch, double-Chinchilla-tokens model. gpjt/jax-with-mha-bias-no-dropout-2-epoch -- the model trained on two epochs over the Chinchilla-optimal number of tokens. "A quick(ish) Chinchilla check" gpjt/jax-with-mha-bias-larger-chinchilla-1 -- the slightly-larger model. gpjt/jax-with-mha-bias-larger-chinchilla-2 -- the slightly-smaller model. I've also added links to the posts in question.

3rd Sep 2026 • 1 votes
A quick(ish) Chinchilla check

I recently overtrained a couple of GPT-2 style models, training them both on 40 tokens per parameter rather than the 20 per parameter that is generally regarded as "Chinchilla-optimal". The normal heuristic is that instead of doing that, you should scale up the number of tokens and the number of parameters equally -- so I would have been better off scaling up the model by 2 and the token count by the same amount. By doing that, I should expect to get a better model in terms of loss on my held-back test set than I did with my 40-tokens-per-parameter models. My training machine poppy wasn't doing anything, so I decided to give that a go. Would the Chinchilla rule-of-thumb hold up? As you might expect, it did. But it was a surprisingly close-run thing, and could conceivably have been in the noise. Let's take a look. The Chinchilla heuristic If you already know all about the Chinchilla paper -- regular readers in particular must be sick and tired of it by now :-) -- then click here to skip this section. In "Training Compute-Optimal Large Language Models", which is always called the Chinchilla paper after the name of the model they trained at the end, the authors tried to work out the optimal number of tokens to train an LLM on based on its number of parameters. In particular, they were pushing back on a trend they were seeing at the time, where people were making models ever-larger, but not increasing the amount of data they were training on. The authors were all at Google DeepMind, and this was the kind of project that only a large lab could do: they trained "over 400 language models ranging from 70 million to over 16 billion parameters on 5 to 500 billion tokens". Their conclusion was "for compute-optimal training, the model size and the number of training tokens should be scaled equally: for every doubling of model size the number of training tokens should also be doubled". They don't actually state an overall optimal number of tokens to train on in the paper, but in table 3 they provide an estimate of the optimal training FLOPs and tokens for models of various sizes, and it's approximately 20 tokens per parameter. That number has become a heuristic, and people talk about a model as being trained for the Chinchilla-optimal number of tokens. Models that were trained on fewer tokens per parameter are referred to as "undertrained", and models that were trained on more as "overtrained". It's worth noting that overtraining a model is not, in itself, a bad thing. If you have a model of a particular size and you continue training it past the Chinchilla-optimal number of tokens, it will -- in general -- get better. The point of the heuristic is that doing that is not the best way to spend whatever budget you have in terms of compute time. You'll get better results, as they say, by scaling the number of tokens and the number of parameters equally. But let's say you're creating a model for specific target hardware -- say, a mobile device. You have a hard restriction on how large the model can be -- the device has only so much RAM to hold it. So it might make sense to overtrain to get a better model. 1 But if you're not so limited in how many parameters you can use, then you should indeed scale the model up, and that's what I wanted to try. How would that work? Scaling the model A week or two back, I was investigating whether I could make my GPT-2 style models better at a specific instruction-following task by overtraining them. The details of that experiment aren't important here, but what it meant was that I had three GPT-2-style models, each of exactly the same size, roughly 163M parameters A Chinchilla-optimal one, which I'll call jax-gpt2-chinchilla here. One trained on twice the Chinchilla-optimal tokens, jax-gpt2-2x-chinchilla. One trained on the Chinchilla-optimal tokens, with two epochs (so that it was trained for as long as #2): jax-gpt2-2-epoch-chinchilla When I tested them against a held-back test set of sequences -- stuff that they'd never seen before -- they got results rather like you might expect: Test loss jax-gpt2-2x-chinchilla 3.324953 jax-gpt2-2-epoch-chinchilla 3.326482 jax-gpt2-chinchilla 3.418784 A lower loss is better, and you can see that the longer-trained models were noticeably better than the Chinchilla-optimal one. The difference between them was tiny; they were trained starting with the same initial weights, and the training runs themselves were deterministic, but a difference of 0.05% in loss doesn't seem like it could be meaningful -- an extra batch for one or one fewer for the other could easily swap them around, you'd think. Now, these models each had 163,009,536 parameters -- they were the small-size model from the GPT-2 paper, modified to not have QKV bias or weight-tying. jax-gpt2-chinchilla had been trained on 3,260,190,720 tokens (rounded up to fit into a round number of full batches), and the other two on 6,520,381,440 tokens each -- double the amount (rounded up too). What I needed to do for my Chinchilla check was to try training a model that used the same amount of compute, scaling the parameters and the number of training tokens equally. Because training compute increases roughly linearly with both parameters and tokens, that would mean scaling both up by 2, giving us: 163,009,536*2≈230,530,296parameters ...and thus 4,610,605,920 tokens. How to scale the model up? In the GPT-2 paper, they train four models: Name Parameters 2 Layers d_emb MHA heads 3 small 124M 12 768 12 medium 345M 24 1024 16 large 762M 36 1280 20 xl 1542M 48 1600 25 I wanted to scale my own model up from 163M parameters to about 231M. Which of those numbers would I want to increase, and by how much? The first thing that stands out is that the number of heads is always 1/64th of the number of embedding dimensions. So that sorted that one out. I just needed to adjust the number of layers, and the number of embedding dimensions, but ensure that the latter was a multiple of 64. I decided to see if I could fit some kind of curve to the relationship between the number of parameters and the GPT-2 authors' choices. This was made a bit more complicated by one thing: they were using weight-tying, and I was not. That meant that they re-used the embedding matrix at the start of the LLM as an output head at the end -- which is why they had 38M fewer parameters. Embeddings and the output head make up a surprisingly large percentage of the parameters for small models like this -- about 47% without weight-tying, 23% with. I couldn't work out a solid way to scale things up and wound up doing some rather messy hacking around in a spreadsheet. I came up with two proposed model sizes that were within a couple of percentage points of the right size: Name Layers d_emb MHA heads Parameters % diff slightly-larger 15 896 14 235,621,120 +2.21% slightly-smaller 14 896 14 225,978,368 -1.97% Interestingly, I found that because d_emb could only change in increments/decrements of 64, it was a pretty coarse control -- my first attempt at making a slightly-smaller model changed it to the next step down, 832, but that led to a model that was 9.25% too small. That was an interesting first lesson. I'd previously been thinking of the Chinchilla rule as being something like "don't double the tokens, just scale the model and the tokens equally". But that "just" was wrong. Scaling a model is hard -- even with just two dials to fiddle with, like in this case, it was tricky to get something right -- and I can't say for sure that my choices were the right ones. Anyway, the next step was to double-check that these models would use the right amount of compute to train. Training FLOPs As I said earlier, the compute time scales roughly linearly with the number of parameters. Let's dig into that "roughly". Different kinds of parameters take different amounts of FLOPs to train, and scale differently with things like the embedding dimensions, sequence length, and so on. Now, for very large models, a lot of that comes out in the wash, but with tiny models like these where the embeddings make up such a large proportion of the parameters, it might matter. Conveniently, in appendix F of the Chinchilla paper, they provide a set of formulae for estimating the number of training FLOPs for a normal dense LLM like these ones. I coded that up into a script that, given the JSON configuration files I was using for my models and training runs, would work out the number of FLOPs for a single epoch of training. It didn't take account of the fact that my real training runs round the number of tokens up so that we do a round number of full batches, but I felt that so long as the results weren't very close that wouldn't matter. I got these results (multiplying the two-epoch numbers by two): Est. training FLOPs jax-gpt2-chinchilla 3,544,967,596,946,227,200 jax-gpt2-2x-chinchilla 7,089,935,193,892,454,400 jax-gpt2-2-epoch-chinchilla 7,089,935,193,892,454,400 slightly-larger 7,419,664,885,127,577,600 slightly-smaller 6,804,429,215,367,168,000 The numbers were indeed different enough that I wasn't worried about the batch-rounding. And the good news was that slightly-larger and slightly-smaller would indeed use slightly more and slightly less compute to train than the overtrained models -- about 4.6% more and 4% less respectively. A true Chinchilla-equivalent model would lie somewhere between them. It was time to train some models! Training I kicked off the run for the slightly-larger model first. Because it was bigger than the 163M models I'd been training, I couldn't fit such large batches into my VRAM; previously I'd been running with a batch size of 6, and now I could only fit in a batch of 4. Luckily, though, I was using gradient accumulation, so by bumping that up from 16 steps to 24 steps I could keep the same overall batch size and keep the training runs comparable. Even despite that, the training run ran out of VRAM about 60 hours in -- I'm guessing due to VRAM fragmentation, as I did not have TF_GPU_ALLOCATOR set to cuda_malloc_async -- but I was able to restart from the most recent checkpoint and complete the run. After just less than four days total training time, it completed. When it was done, I copied the last checkpoint 4 over to my dev box, perry, and ran my standard smoke test against it, asking it to complete "Every effort moves you" with 20 tokens, using greedy sampling. I got something reasonably coherent: Every effort moves you. I’m not sure what you’re thinking. I’m Next, I converted the safetensors file -- which had been saved by my JAX code -- into a format compatible with my PyTorch code, because that's what I use for evals. I ran another smoke test (this one with temperature 1): Every effort moves you through the motions for your life, your soul, your body, and your soul’s happiness Very spiritual. Next, it was time to work out the loss on my held-back test set: 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-larger-chinchilla/model.json ../jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-larger-chinchilla/checkpoints/latest/pytorch-model.safetensors Fetching 4 files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 4/4 [00:00<00:00, 3485.09it/s] 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 3200/3200 [07:11<00:00, 7.42it/s] Loss against our test dataset: 3.280028 Well, it was certainly better than the 3.324953 that the best of the overtrained models got -- but only by a bit over 1% better. Interesting! I decided to train the second model, slightly-smaller. This one crashed mid-way through with an error that I've seen before: 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'] I'm going to have to investigate that more in future, but for now, I just restarted from the checkpoint, and again after a bit less than four days, I had a model. The JAX smoke test was solid: Every effort moves you forward. The best way to get started is to start with a free trial. You can ...and so was the PyTorch one: Every effort moves you forward in love with our products. I love the way it’s easy to use. Both quite commercial this time! It was time for the proper test loss eval: 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-larger-chinchilla-2/model.json ~/Dev/jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-larger-chinchilla-2/checkpoints/latest/pytorch-model.safetensors Fetching 4 files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 4/4 [00:00<00:00, 3151.24it/s] 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 3200/3200 [06:45<00:00, 7.90it/s] Loss against our test dataset: 3.292937 So, slightly worse than the 3.280028 from the larger model, better than the 3.324953 from the best overtrained one. Time to put this all together. Results Here's an updated version of the table from the start of this post; I've added in the two new models, and the improvement they each had over jax-gpt2-2x-chinchilla in both absolute terms and as a percentage rounded to 3sf. Test loss Improvement Improvement % slightly-larger 3.280028 0.044925 1.35% slightly-smaller 3.292937 0.032016 0.962% jax-gpt2-2x-chinchilla 3.324953 - - jax-gpt2-2-epoch-chinchilla 3.326482 - - jax-gpt2-chinchilla 3.418784 - - Now, unlike the overtrained models, prior to training these two new ones started with different initial weights to the jax-gpt2-chinchilla one -- after all, they had to, because they had more of them! A while back, I did a bit of analysis of how random variation in weight initialisation can change the resulting test loss. It wasn't anything in-depth, but I trained three models with different explicit seeds set prior to the model initialisation, but with the same seed set before the training run started 5. Those three models wound up with test losses of 3.681356, 3.673943, and 3.664345. Doing statistics with three data points is a bit flaky, but the cost of training models is so high that I'll leave the Proper Science to the likes of Google DeepMind and wing it :-) Mean: ~3.673215 Sample variance: ~0.000073 Standard deviation (SD): ~0.008529 Now, piling statistical flakiness on statistical flakiness, we'll compare these. You'd normally expect about two thirds of results to be within one SD of the mean, 95.4% to be within two SDs, and 99.7% to be within three. Three SDs on that (yes, different, I know) distribution is 0.025587. That's smaller than both of the improvements that our Chinchilla-optimal runs had over the overtrained ones. So what does that tell us? Well, perhaps not much given the statistical flakiness. But I think it is useful directionally. It suggests that we might be able to take these results seriously as an improvement, and that Chinchilla held: scaling up the model and the number of tokens evenly did give us a better model than just scaling up the number of tokens. In particular, the fact that the loss for slightly-smaller was lower -- even though it had 4% less compute spent on it than the overtrained models -- was encouraging. But it's certainly far from a slam-dunk. A larger test, training lots of overtrained models and lots of Chinchilla-optimal ones, all with different random seeds, would give actual real serious data. Not worth it for me, and perhaps not for anyone. Conclusion I wanted to do a quick sanity check of the Chinchilla heuristic of 20 tokens per parameter. I came up with results that were certainly in line with it -- perfectly so in terms of the ordering of the models I trained. But the effect was small enough that I could imagine that it was in the noise, especially given the small numbers of models I'm able to train. I'll chalk it up as a tentative success. In addition, I learned one useful thing: when talking about scaling up a model to more parameters, you actually have to think quite hard about where you want to put those parameters. I wound up doing a rough curve-fit to the models in the GPT-2 paper, but I have no idea if that was optimal. At some point I should try to dig up some research into optimising embedding dimensions, numbers of layers, and so on. But not now, as I've a bunch of other stuff I want to investigate first. Anyway, I hope you found this experiment interesting, and as ever, comments and questions welcome below. Thanks for reading! I'm less familiar with arguments for under-training -- that is, for fewer than 20 tokens per parameter. I've heard that these days, modern LLMs get a lot more reinforcement learning than they do pre-training, and perhaps that might mean that some very big ones are undertrained prior to RL? I'm uncertain. It's unlikely to be raw lack of data; even for those of us outside the big labs, FineWeb has 18.5T tokens. On its own, that would be enough to train a 0.925T-parameter model, and given that you can apparently do four epochs over the same data before you start getting diminishing returns, that takes us up to 3.7T. That's frontier-lab size, and I'm sure they have better datasets than FineWeb. ↩ Parameter counts are from the paper, apart from the "small" model, which is known to be wrong -- I used my own calculation, and the result is in line with what I've seen elsewhere. ↩ The paper doesn't mention the number of heads; these numbers are from "Build a Large Language Model (from Scratch)", and match up with the ones on this Hugging Face page. ↩ Regular readers might have noticed that I'm ignoring what I've been calling the "best" checkpoint. I've come to the conclusion that because for my training script, "best" means best in terms of training loss, and the training loss changes based on what training data the model has seen recently, it's actually not a very useful metric and just confuses things. At some point I'll probably re-introduce pre-checkpoint evals and use that for "best", which would be the right way to do it. ↩ At the time I was using dropout, so training runs were not deterministic without a known seed. ↩

7th Aug 2026 • 2 votes
Writing an LLM from scratch, part 34b -- from bigrams to GPT-2, one component at a time (in JAX)

This post is the capstone of the most long-running series on my blog. In December 2024 (!), I started reading Sebastian Raschka's book "Build a Large Language Model (from Scratch)", and worked through it carefully. Being who I am, despite trying to apply a strict "no side quests" policy, I found myself zooming off and digging into all kinds of things. It's time to wrap it up. I had decided that the endpoint would be to build and train an LLM from scratch just using my notes -- no reference to the book, no reference to the model code I'd written when following the book. After an X/Twitter poll, I decided to use JAX for that, just to make sure that I really was building it from scratch and not regurgitating bits of PyTorch code like a bad coding LLM spitting out half-digested lumps of Stack Overflow. In my last post, I showed how I built a JAX training script that mirrored what I had built for the original PyTorch version of the model. To test it as I went along, I used it to train a really dumb "LLM", which instead of trying to predict the next token for every token in an input sequence, instead predicted the input -- that is, if you fed it The fat cat sat on the mat It would return the same thing. I called that an A-to-A model. In this post, I'll show you how I turned it into a GPT-2 model, and then trained it from scratch on my RTX 3090 (using the parameter counts for the original paper's "small" size). What turned out really well with this is that I found a route that meant that almost every component I added made the model better! That's not guaranteed -- sometimes different aspects of an AI model depend on each other, so adding A without also adding B makes things worse. But (admittedly with a bit of backtracking in places) I was able to find a route that shows a nice clear progression. The final training run took 37 hours 15 minutes -- compared to 40 hours, 38 minutes for an equivalent PyTorch model. That is despite it being full-fat 32-bit -- the PyTorch one was using Automatic Mixed Precision (AMP), which allowed it to use 16-bit calculations in places where it would be relatively harmless in terms of loss. When asked to continue "Every effort moves you", it came back with a decent response: Every effort moves you closer to your goals, but if you are unsure of what it takes, you don’t The model got 3.418784 loss on my held-back test dataset, as compared to my PyTorch model's 3.538161, and even more impressively, it was better than the original GPT-2 small's result of 3.499677 on the same dataset! However, just as I found previously, the OpenAI weights still beat mine consistently in instruction fine-tuning challenges. Let's get started. The starting point -- A-to-A At the end of the last post, we had a solid training loop, using all of the tricks I'd picked up with my PyTorch code. The A-to-A model we were training with it looked like this: from flax import nnx class GPTModel(nnx.Module): def __init__( self, vocab_size, context_length, d_emb, n_heads, n_layers, qkv_bias, drop_rate, rngs, ): self.token_embedding = nnx.Embed( num_embeddings=vocab_size, features=d_emb, rngs=rngs, ) self.output_head = nnx.Linear( in_features=d_emb, out_features=vocab_size, use_bias=False, rngs=rngs, ) def __call__(self, xs): input_embeddings = self.token_embedding(xs) return self.output_head(input_embeddings) That was based on my preferred model of how LLMs work, where at the top level for a model, we feed in a sequence of token IDs, then: Firstly, we convert them into embeddings, so we get a sequence of vectors, one for each token. We do this by a lookup into a table, but we can see it conceptually as a projection via a matrix, from vocab space (where a particular token ID is a one-hot vector) to embedding space. Next, we do the magic with our Transformers layers, getting embeddings for the next token. After these layers, the embedding at position n in the output sequence is for the predicted token to come after the token at position n in the input sequence, considering that input token and all other tokens to its left. Finally, we project those back from embedding space to logits, this time actually using a real matrix (in the form of a linear layer), the output head. The logits (after being run through softmax) represent the probabilities for each token of it being the next one. The A-to-A model basically skipped the second step completely: it would project to embedding space, then immediately project back to vocab space -- and after training, it was pretty good at mapping a sequence to itself. One interesting question is, if we train the same code, but this time try to get it to make next-token predictions, how good will it be at that? Obviously it can't be as good as a full LLM. But there are correlations between tokens; full stops will generally be followed by spaces, adjectives will normally be followed by other adjectives or nouns (at least in English), and so on. It would be kind of like the predictive text systems on a phone, where (at least until recently) it would just use the last word you entered to generate a list of possible next words to select from. Old-school natural language processing has a name for this: bigrams. The idea is that you can work out statistically what the most common two-word pairs are, which allows you to make a guess at a next word from a single one. (There are also trigrams, where you look at the last two words when predicting the next, then 4-grams, 5-grams, and so on.) You'd build up a full probability table -- for every word in your vocab, you'd have the probability of every word coming next. So maybe even with that minimal model, we could get it to learn something similar to a set of token-level (rather than word-level) bigrams, which would then get the loss down. Obviously it wouldn't be as good as a full bigram table -- for our GPT-2 vocab size of 50,257, that would need 50,2572=2,525,766,049 parameters -- but perhaps it could approximate one. (For comparison, the model we're using has just an embedding table and an output head, each mapping between 50,257 dimensions and 768, so that's 2×50,257×768≈77 million parameters -- about 3% of the full table.) An uninitialised model would (hopefully) have a loss of about 10.82, implying a perplexity equal to the vocab size. If we can train our dumb model to get better loss than that, then we'd have the beginnings of an LLM. That was a simple test to run. In my training code, I had a dataset class that looked like this: class BigTrainDataset: def __init__(self, all_tokens, seq_length, microbatch_size): self.xs = all_tokens[:-1].reshape(-1, microbatch_size, seq_length) self.ys = all_tokens[:-1].reshape(-1, microbatch_size, seq_length) def __getitem__(self, ix): return self.xs[ix], self.ys[ix] def __len__(self): return self.xs.shape[0] That is, the inputs, the xs, were the same as the targets, the ys. If we fed it The fat cat sat on the mat ...then we'd be training it to output exactly the same thing. The modified version for a real LLM would involve feeding it something like this: The fat cat sat on the ...and targeting this: fat cat sat on the mat That's a simple change -- that __init__ method became this: self.xs = all_tokens[:-1].reshape(-1, microbatch_size, seq_length) self.ys = all_tokens[1:].reshape(-1, microbatch_size, seq_length) I did that, and kicked it off to train on the 92,209,152 tokens that I was (somewhat arbitrarily) using in the last post to test my training loop. The loss chart looked like this: That was pretty promising! Loss came down from roughly 10.82 down to a fairly stable 6 or so by global step 768, and seemed to flatten out there. It's possible that further training could have got it down a bit more, but I decided (again, somewhat arbitrarily) to use the average train loss in the checkpoint period ending at step 937 as my starting point. If we could make changes that reduced that, then we'd be moving forward. For this model, that value was 5.909. So, what were the changes we needed to make to change our bigram-style model to a real, if small, LLM? Building GPT-2: a checklist Adapting from my how LLMs work post, a GPT-2-style LLM looks like this. We receive our sequence of token IDs, and then: Convert them into embeddings. ✔ ️done Add on position embeddings. Run these embeddings through multiple successive Transformers blocks. Layer normalisation Project them back from embedding space to vocab space. ✔ ️done Inside the Transformers blocks, we: Take a copy of the input sequence of embeddings Layer normalisation Run multi-head attention Add the copy back in so that the version that came out of MHA is something more like an "annotation" of the original Take a second copy of that one Layer normalisation again Run it through a simple neural network Add the results of that back in. So that gave me the checklist; looking at it, the most tempting next step was layer normalisation (henceforth LayerNorm). It's used at the end of the core loop, and then twice in the Transformers blocks. What would happen if we coded it up, and then added it to the core only? LayerNorm The purpose of LayerNorm is to stabilise training. We constrain the values flowing through our model so that they have certain statistical properties that tend to make the whole thing more trainable. That would mean that if it did help with this model -- placed in between the embedding layer at the start, and the output head at the end -- then we'd hope for loss to go down faster, and ideally finish at a lower level. Time to code it up! NNX has its own LayerNorm implementation, of course (as does PyTorch), but in the book, we implement it ourselves, and that felt like the correct path to take. Firstly, I implemented a dummy version: class LayerNorm(nnx.Module): def __init__(self): ... def __call__(self, xs): return xs ...and updated the core GPTModel to create and call one: class GPTModel(nnx.Module): def __init__( ... ): self.token_embedding = nnx.Embed( ... ) self.output_norm = LayerNorm() self.output_head = nnx.Linear( ... ) def __call__(self, xs): input_embeddings = self.token_embedding(xs) normalised = self.output_norm(input_embeddings) return self.output_head(normalised) And kicked off a training run for a few seconds just to make sure that it hadn't broken anything and that loss dropped -- being my first NNX module-inside-a-module, I worried that there might have been something non-intuitive that I had to do to get it to work. But everything seemed good -- loss was dropping, no errors. So, following the notes I made when I first learned about LayerNorm, I needed to make the values flowing through centred around zero by subtracting their mean, and then scale them to have a variance of one by dividing by the standard deviation (details in those notes). The shape of the xs I had coming into my LayerNorm class's _call was this: xs.shape=(6, 1024, 768) That was (batch_size, seq_len, d_emb). So we needed to do those operations strictly on the last axis, manipulating each embedding independently. JAX has a std function and a mean one, both of which take an axis parameter. The Array object repackaged those as methods, which was convenient, so I did a first cut test like this: class LayerNorm(nnx.Module): def __init__(self): ... def __call__(self, xs): jax.debug.print(f"{xs.shape=}") means = xs.mean(axis=-1) jax.debug.print(f"{means.shape=}") stds = xs.std(axis=-1) jax.debug.print(f"{stds.shape=}") return xs That printed out these results: xs.shape=(6, 1024, 768) means.shape=(6, 1024) stds.shape=(6, 1024) ...which looked plausible; one number for each embedding vector. Could we broadcast them across the array? class LayerNorm(nnx.Module): def __init__(self): ... def __call__(self, xs): jax.debug.print(f"{xs.shape=}") means = xs.mean(axis=-1) jax.debug.print(f"{means.shape=}") stds = xs.std(axis=-1) jax.debug.print(f"{stds.shape=}") normalized = (xs - means) / stds jax.debug.print(f"{normalized.shape=}") return normalized This blew up: ValueError: Incompatible shapes for broadcasting: shapes=[(6, 1024, 768), (6, 1024)] Fair enough. But mean and std have a keepdims kwarg that looked like it would help: class LayerNorm(nnx.Module): def __init__(self): ... def __call__(self, xs): jax.debug.print(f"{xs.shape=}") means = xs.mean(axis=-1, keepdims=True) jax.debug.print(f"{means.shape=}") stds = xs.std(axis=-1, keepdims=True) jax.debug.print(f"{stds.shape=}") normalized = (xs - means) / stds jax.debug.print(f"{normalized.shape=}") return normalized ...and it did! xs.shape=(6, 1024, 768) means.shape=(6, 1024, 1) stds.shape=(6, 1024, 1) normalized.shape=(6, 1024, 768) Excellent. So the next step was to see if that would work even slightly. Interestingly loss started off a bit higher at 11.29 after the first global step -- so adding in the LayerNorm had actually made the model worse than it was -- but it seemed to be falling rapidly. Things weren't totally broken, at least. But there was more to LayerNorm than just zeroing the mean and scaling to the variance; we also needed to scale them up by a learnable amount, and then shift/bias them by adding on a different trainable amount. More precisely, both of those trainable amounts were different for each of the (in this case) 768 embedding dimensions. We needed two learnable vectors of length d_emb. I hadn't noted it down at the time but I figured (as it turned out, correctly) that a sensible starting point for those values would be all-zero for the bias, and all-one for the scale. From this help page, the way you create a trainable array associated with an NNX module is this: nnx.Param(jax.random.normal(rngs.param(), (dim, dim))) That code created a random vector, rather than the zeros/ones we needed, and we'd need to get the dimensions right. Because of the "Incompatible shapes for broadcasting" error I'd just had, I was feeling a bit paranoid about the latter, so I chose a shape of (1, 1, d_emb), and wrote this: class LayerNorm(nnx.Module): def __init__(self, d_emb): self.scale = nnx.Param(jnp.ones((1, 1, d_emb))) self.bias = nnx.Param(jnp.zeros((1, 1, d_emb))) def __call__(self, xs): jax.debug.print(f"{xs.shape=}") means = xs.mean(axis=-1, keepdims=True) jax.debug.print(f"{means.shape=}") stds = xs.std(axis=-1, keepdims=True) jax.debug.print(f"{stds.shape=}") normalized = (xs - means) / stds jax.debug.print(f"{normalized.shape=}") scaled_and_biased = (normalized * self.scale) + self.bias return scaled_and_biased That looked pretty plausible, though in retrospect I think I was being overly cautious and didn't need the leading two axes for the scale and bias. The only thing I was unsure about was whether the nnx.Param wrappers I had put in were really making those arrays trainable. I put some code in to print them out and kicked off a run for a few minutes, and confirmed that they were changing in ways that seem plausible -- small non-zero bias, scale close to but not equal to one. That was all good! Next, I spotted one issue. What if one of the standard deviations was zero? That would lead to a divide-by-zero error here: normalized = (xs - means) / stds Now, the standard deviation, if it's not zero, has to be positive -- so adding on a small value would fix that 1: normalized = (xs - means) / (stds + 1e-5) With that in place, I felt that it was ready to go. Time to do a full training run! I kicked that off, and it completed with this output: 2026-06-20 19:08:17.189721 Tokens seen: 92,209,152 2026-06-20 19:08:17.189724 Throughput: 95,383 tokens/second 2026-06-20 19:08:17.189734 Final train loss: 5.736 2026-06-20 19:08:17.189737 Done Loss looked like this: Let's look at the results for the previous run without LayerNorm for comparison: You can see that the new run, the first one, drops faster. It's harder to see from the chart, but it also finished up with a lower training loss at 937 (my relatively arbitrary metric): 5.734 rather than 5.909. That was interesting! The new model was basically doing the same thing -- predicting the next token based only on the "current" token, but loss was lower. My take is that if we had trained the non-LayerNorm model for longer, it might have managed to eventually grind out a better loss. But LayerNorm was doing its job -- it was stabilising training, and as a result we converged faster. That was a win! I decided to run it through my old smoke test from the PyTorch training runs, and see how it completed "Every effort moves you": Every effort moves you can be a few years. -year-year-year-year-year-year- It was kind of impressive that it managed to finish the first line before it got stuck in a loop -- but it was understandable that we couldn't expect anything good yet. Each predicted token was based entirely on the token before it. What next? Back to our checklist: Convert token IDs into embeddings. ✔ ️done Add on position embeddings. Run these embeddings through multiple successive Transformers blocks. Layer normalisation ✔ ️done Project them back from embedding space to vocab space. ✔ ️done Inside the Transformers blocks, we: Take a copy of the input sequence of embeddings Layer normalisation Run multi-head attention Add the copy back in so that the version that came out of MHA is something more like an "annotation" of the original Take a second copy of that one Layer normalisation again Run it through a simple neural network Add the results of that back in. So, at this stage, for each input token we were predicting the next one based on the input token only -- like I said earlier, we were doing a somewhat roundabout way of building an approximation of a table of bigram probabilities. What would happen if we started paying attention to the tokens to the left? And what would be the simplest, dumbest way to do that? A single layer of single-head attention The real LLM has multiple layers of multi-head attention, each one also having a feed-forward network, some LayerNorms, and some shortcut connections. Single-head attention is easier to code, but even on its own, you'd expect it to be able to add some value. Each token would get at least some information from the ones to the left. And one layer, likewise, you'd expect might help a bit. I suspected that it wouldn't work on its own -- I expected I'd need shortcut connections too -- but decided to start with attention on its own. I modified the main class to have a single "Transformers" layer: class GPTModel(nnx.Module): def __init__( ... ): self.token_embedding = nnx.Embed( ... ) self.transformers_layer = TransformersLayer(d_emb, qkv_bias, rngs) self.output_norm = LayerNorm(d_emb) self.output_head = nnx.Linear( ... ) def __call__(self, xs): input_embeddings = self.token_embedding(xs) transformed = self.transformers_layer(input_embeddings) normalized = self.output_norm(transformed) return self.output_head(normalized) ...where that layer was actually just single-head attention: class TransformersLayer(nnx.Module): def __init__(self, d_emb, qkv_bias, rngs): self.attention = Attention(d_emb, qkv_bias, rngs) def __call__(self, xs): return self.attention(xs) Next, it was time for the Attention class. I'm not going to write yet another attention explainer -- I think my "How do LLMs work?" one does a decent job of that, and "The 'why' of attention, or: attention heads are dumb" works well too. So in the next bit I'll assume that you understand the basics. My first cut was basically just the maths (up to the causal mask) to get the attention scores: class Attention(nnx.Module): def __init__(self, d_emb, qkv_bias, rngs): self.d_emb = d_emb self.W_q = nnx.Linear(d_emb, d_emb, use_bias=qkv_bias, rngs=rngs) self.W_k = nnx.Linear(d_emb, d_emb, use_bias=qkv_bias, rngs=rngs) self.W_v = nnx.Linear(d_emb, d_emb, use_bias=qkv_bias, rngs=rngs) def __call__(self, xs): Q = self.W_q(xs) K = self.W_k(xs) V = self.W_v(xs) omega = Q @ K.T omega /= jnp.sqrt(self.d_emb) causal_omega = jnp.tril(omega) It did the projections into query, key and value space, worked out the attention scores with the array multiplication, normalised it by dividing by the square root of the number of dimensions in the Q-K embedding space, and then zeroed out the scores where a token was attending to tokens in its "future". There were a couple of problems, though. Firstly, that wouldn't work if we were working with batches, and secondly, zeroing out the non-causal scores wasn't quite correct. The batches first. Our incoming xs here would have the shape (batch_length, seq_len, d_emb). After the projections to the Q-K embedding space, both Q and K would also be shaped (batch_length, seq_len, d_emb). Now, the .T property on the JAX array class just reverses the axes, so the code above would give us K.T with the shape (d_emb, seq_len, batch_length). That would break! Matrix multiplication in JAX expects all but the last two axes to represent batches, so we actually wanted K.T to have the shape `(batch_length, d_emb, seq_len). That meant that what we actually wanted was to just transpose the last two axes. The JAX transpose function takes an axes parameter that allows you to specify the specific re-ordering of the input axes that you want. So I could rewrite the code like this: def __call__(self, xs): Q = self.W_q(xs) K = self.W_k(xs) V = self.W_v(xs) omega = Q @ jnp.transpose(K, axes=(0, 2, 1)) omega /= jnp.sqrt(self.d_emb) causal_omega = jnp.tril(omega) As Q would have the shape (batch_length, seq_len, d_emb), and the transposed version of K would be (batch_length, d_emb, seq_len), they'd be compatible for matrix multiplication and give us a result that was (batch_length, seq_len, seq_len) -- just what we wanted for attention scores. The next step was to fix the causal mask. The next step in this attention mechanism was going to be running the causal attention scores in causal_omega through softmax over the last dimension, to convert them into attention weights. Now, our current code was zeroing out unwanted acausal scores, but a zero still contributes to softmax. If you want a particular value to come out of softmax guaranteed to be zero, you need to set it to minus infinity. I decided that the easiest way to do this was to create a causal mask -- a boolean array that matched the size of omega, but was full of Trues: causal_mask = jnp.ones_like(omega, dtype=bool) Then I could zero out (well, "false out") the cells in the mask related to unwanted future-facing scores, just like I was previously doing on the scores: causal_mask = jnp.tril(causal_mask) ...and then I could apply that mask to omega with jnp.where, telling it to create a new array, taking the value from omega where the mask had True, and -jnp.inf in places where it had False. causal_omega = jnp.where(causal_mask, omega, -jnp.inf) That seemed solid, so I just needed to run the result through jax.nn.softmax, specifying that the last dimension was the one where it should apply the function, and that would give me the attention weights: attention_weights = jax.nn.softmax(causal_omega, axis=-1) Finally, I just needed to use those attention weights to get the attention output by mixing in appropriate portions of the projection of the inputs into value space, V: return attention_weights @ V As attention_weights was shaped (batch_length, seq_len, seq_len), and V (like Q and K) was shaped (batch_length, seq_len, d_emb), the batch axes were at the start where they belonged, and the matrix multiplication would work and return something shaped (batch_length, seq_len, d_emb). With that, we were done! The final single-head attention class looked like this: class Attention(nnx.Module): def __init__(self, d_emb, qkv_bias, rngs): self.d_emb = d_emb self.W_q = nnx.Linear(d_emb, d_emb, use_bias=qkv_bias, rngs=rngs) self.W_k = nnx.Linear(d_emb, d_emb, use_bias=qkv_bias, rngs=rngs) self.W_v = nnx.Linear(d_emb, d_emb, use_bias=qkv_bias, rngs=rngs) def __call__(self, xs): Q = self.W_q(xs) K = self.W_k(xs) V = self.W_v(xs) omega = Q @ jnp.transpose(K, axes=(0, 2, 1)) omega /= jnp.sqrt(self.d_emb) causal_mask = jnp.ones_like(omega, dtype=bool) causal_mask = jnp.tril(causal_mask) causal_omega = jnp.where(causal_mask, omega, -jnp.inf) attention_weights = jax.nn.softmax(causal_omega, axis=-1) return attention_weights @ V I kicked off a training run with that, and it did work, in that loss went down over the course of the run -- but at the end of the run, the loss at step 937 was 5.934 -- significantly above the 5.734 I got on the previous run, with no attention. But that made sense! As I'd said earlier, I suspected that this wouldn't help if we had no shortcut connection. Intuitively, if you want to work out what token should be at position n+1, on average the most important other token you need to know about is probably whichever one is at position n. Knowing about the tokens at n−1, n−2, and so on, could well be helpful -- maybe very helpful -- but not at the cost of not knowing about the one at n. Now, single attention heads are just simple pattern-matchers. They can't learn complex rules, it's only by working together -- "horizontally", in multi-head attention or "vertically" across multiple layers -- that they can do complex things. What we were asking this head to do was to learn some way of gathering information about previous tokens, and also to keep the knowledge about the "current" one. That's a tall order for a dumb attention head! In my mind, this is a large part of the benefit of shortcut connections. They are often presented as a way to make sure that during training, gradients flow smoothly from the output end of the model to the earlier layers. But I prefer to think of them as preserving the original embeddings, so that each layer doesn't completely replace what came into it, but instead does something closer to adding on its own notes -- like scholars adding commentary to a core text in the Talmud. In the training run above, the attention head was trying to learn how to preserve the meaning of the embedding it was working on, while also merging in information from earlier ones. If we added a shortcut connection, then it would only have to do the second of those two jobs. The code was simple: I updated the TransformersLayer module to do a shortcut connection: class TransformersLayer(nnx.Module): def __init__(self, d_emb, qkv_bias, rngs): self.attention = Attention(d_emb, qkv_bias, rngs) def __call__(self, xs): shortcut = xs att = self.attention(xs) return shortcut + att I kicked off a training run, and at the end it printed this: 2026-06-23 03:51:18.086097 Tokens seen: 92,209,152 2026-06-23 03:51:18.086099 Throughput: 90,442 tokens/second 2026-06-23 03:51:18.086108 Final train loss: 5.570 2026-06-23 03:51:18.086121 Done The loss chart looked like this: And, importantly, that training loss at step 937 which I was using as a metric was 5.553 -- a decent improvement over the previous best of 5.734. Even a dumb single attention head was able to do something useful, if it had a shortcut connection. I decided to run another qualitative smoke test: Every effort moves you can be able to get to get to get to get to get a lot of the way to get I mean, it was repetitive, but it was actually getting noticeably closer to making sense! So that was excellent news. What next? Our checklist looked like this: Convert token IDs into embeddings. ✔ ️done Add on position embeddings. Run these embeddings through multiple successive Transformers blocks. part-done -- one layer only Layer normalisation ✔ ️done Project them back from embedding space to vocab space. ✔ ️done Inside the Transformers blocks, we: Take a copy of the input sequence of embeddings ✔ ️done Layer normalisation Run multi-head attention part-done -- single-head attention only Add the copy back in so that the version that came out of MHA is something more like an "annotation" of the original ✔ ️done Take a second copy of that one Layer normalisation again Run it through a simple neural network Add the results of that back in. Now, our single attention layer was lacking something. Without position embeddings, that layer has no idea what order the tokens before the one it's looking at come in. If it's considering the " cat" in The fat cat ...it doesn't know if it's looking at "The fat cat" or "fat The cat". Position embeddings are simple, and might help, so that was the next step. Position embeddings These were trivial to add. We had this core code: class GPTModel(nnx.Module): def __init__( ... ): self.token_embedding = nnx.Embed( ... ) self.transformers_layer = TransformersLayer(d_emb, qkv_bias, rngs) self.output_norm = LayerNorm(d_emb) self.output_head = nnx.Linear( ... ) def __call__(self, xs): input_embeddings = self.token_embedding(xs) transformed = self.transformers_layer(input_embeddings) normalized = self.output_norm(transformed) return self.output_head(normalized) So I just added a position encoding module in __init__: self.position_embedding = nnx.Embed( num_embeddings=context_length, features=d_emb, rngs=rngs, ) ...and mixed it in with the token embeddings to create new, improved input_embeddings to be used in our "Transformers" layer: token_embeddings = self.token_embedding(xs) b, n = xs.shape position_embeddings = self.position_embedding(jnp.arange(n)) input_embeddings = token_embeddings + position_embeddings I kicked off a training run with that: 2026-06-23 04:44:44.759768 Tokens seen: 92,209,152 2026-06-23 04:44:44.759771 Throughput: 88,618 tokens/second 2026-06-23 04:44:44.759779 Final train loss: 5.386 2026-06-23 04:44:44.759781 Done Pretty hard to distinguish from the previous one, but the metric I was tracking, that loss at step 937, had improved again! We were down to 5.354 from 5.553 :-) A quick qualitative smoke test didn't show that improvement, though: Every effort moves you can be able to get to get to get back to get back to get back to get back to Pretty much indistinguishable to the previous one. But still, Loss Number Went Down, and that's what was important at this stage. It was time to try the next step. From the checklist: Convert token IDs into embeddings. ✔ ️done Add on position embeddings. ✔ ️done Run these embeddings through multiple successive Transformers blocks. part-done -- one layer only Layer normalisation ✔ ️done Project them back from embedding space to vocab space. ✔ ️done Inside the Transformers blocks, we: Take a copy of the input sequence of embeddings ✔ ️done Layer normalisation Run multi-head attention part-done -- single-head attention only Add the copy back in so that the version that came out of MHA is something more like an "annotation" of the original ✔ ️done Take a second copy of that one Layer normalisation again Run it through a simple neural network Add the results of that back in. We had only one attention head right now. Individually, attention heads are dumb, so switching to multi-head attention seemed like a good thread to pull. Multi-head attention At this point, my single-head attention code looked like this: Q = self.W_q(xs) K = self.W_k(xs) V = self.W_v(xs) omega = Q @ jnp.transpose(K, axes=(0, 2, 1)) omega /= jnp.sqrt(self.d_emb) causal_mask = jnp.ones_like(omega, dtype=bool) causal_mask = jnp.tril(causal_mask) causal_omega = jnp.where(causal_mask, omega, -jnp.inf) attention_weights = jax.nn.softmax(causal_omega, axis=-1) return attention_weights @ V I decided to re-implement multi-head attention (which I'll call MHA from here onwards) from first principles rather than working strictly from my notes, and then to come back and check it. If you're looking at your browser's scrollbar with horror ("still only 50%?!") and really don't want to read a full derivation of MHA, you can skip straight to the first complete version of the code. The point of MHA is that we're running multiple copies of the calculation above in parallel -- let's pin down the name of the number of copies as n_heads. Now, we could naively implement it just by spinning off n_heads threads and running the existing code in each, but that wouldn't really take advantage of the GPU's inherent parallelism. I felt that we could rely on the fact that JAX's matrix multiplications treat all but the last two dimensions as "batches". For example, if you have two arrays with shapes: (a, b, c, ..., l, m, n) and (a, b, c, ..., l, n, p) ...then you can multiply them. A m×n matrix multiplied by a n×p one will be m×p, so you'll get something that is (a, b, c, ..., l, m, p) The other dimensions (so long as they match) will essentially act as an a×b×c×...×l batch. Now, right now we were just using a single batch dimension. Let's look at the core multiplication in the attention mechanism, which works out omega, the attention scores. I had this: omega = Q @ jnp.transpose(K, axes=(0, 2, 1)) Breaking that apart into two steps: K_transpose = jnp.transpose(K, axes=(0, 2, 1)) omega = Q @ K_transpose We got K from this line: K = self.W_k(xs) Let's look at the shapes here. xs is our input embeddings for this layer; its shape is (batch_size, seq_len, d_emb). Projecting it through W_k, which is shaped (d_emb, d_emb) gives us a shape for K of (batch_size, seq_len, d_emb) again. Q, being a projection of xs through W_q, which is the same shape as W_k, will have the same shape as K. Now, that means that K_transpose is (batch_size, d_emb, seq_len), and the calculation omega = Q @ K_transpose ...is doing a batched matrix multiplication getting us the omega that we want, shaped (batch_size, seq_len, seq_len). But as I said above, there's no need to stop with just one batch dimension. Let's say that we have n_heads heads, and that they each work with embeddings sized d_head. Imagine that we've already somehow done multiple projections into the key and query spaces for each of our n_heads heads, and that the results have somehow been put into arrays such that Q and K are shaped (batch_size, n_heads, seq_len, d_head) -- that is, we've gained an extra axis that keeps the projections for each head into its query-key space separate. We could use the fact that both of those two leading axes are basically just batch dimensions, and the existing single matrix multiplication will still work, with one tiny tweak: the current transpose is this: K_transpose = jnp.transpose(K, axes=(0, 2, 1)) omega = Q @ K_transpose ...to swap around the last two axes of a three-axis array. With one extra batch dimension, we'll need to take account of that and do this instead: K_transpose = jnp.transpose(K, axes=(0, 1, 3, 2)) omega = Q @ K_transpose That will be a multiplication of Q, shaped (batch_size, n_heads, seq_len, d_head), with K_transpose, shaped (batch_size, n_heads, d_head, seq_len), which gives us an omega of the right shape, (batch_size, n_heads, seq_len, seq_len). So, if we can start treating the heads as just another batch dimension, things seem simpler, at least for the attention score calculation. Let's continue down through the single-head code, and then come back later to how we might get the inputs into that double-batched shape. The next line after the omega calculation just scales the attention scores by a scalar: omega /= jnp.sqrt(self.d_emb) That looked fine, just a broadcast division-by-float. We'd need to change that self.d_emb to be d_head in some manner, but that's all. Next: causal_mask = jnp.ones_like(omega, dtype=bool) The jnp.ones_like will give us an array that's (batch_size, n_heads, seq_len, seq_len) full of Trues. That seems reasonable. The next step: causal_mask = jnp.tril(causal_mask) What will that do? Well, per the tril documentation: When m.ndim > 2, jnp.tril operates batch-wise on the trailing axes. ...which sounded good. batch_size and n_heads would be treated as batch axes, which meant that the next line: causal_omega = jnp.where(causal_mask, omega, -jnp.inf) ...would work. Likewise, with the next line: attention_weights = jax.nn.softmax(causal_omega, axis=-1) ...the axis to apply softmax to is explicitly stated as the last one, which is what we wanted. So at the end of all of those steps, we'd have attention_weights shaped (batch_size, n_heads, seq_len, seq_len), where the last axis had been softmaxed (softmaxxed?). The next line looked a little trickier: return attention_weights @ V In the single-head version we had attention_weights of shape (batch_size, seq_len, seq_len), and V of shape (batch_size, seq_len, d_emb), so multiplying them gives us (batch_size, seq_len, d_emb) In the new MHA code so far, we had our attention_weights shaped (batch_size, n_heads, seq_len, seq_len). So in order for the matrix multiplication to work, we'd need V to be shaped (batch_size, n_heads, seq_len, d_head). That would give us a result shaped as (batch_size, n_heads, seq_len, d_head). And conveniently, we'd already decided that the correct shape for Q and for K was (batch_size, n_heads, seq_len, d_head). If we could use the same "magic" to do the projection into value space -- that is, to get V such that the heads formed a new batch-like axis like we had for Q and K -- then we'd be all set. So, at that point, I'd worked out the core of MHA. If we could get all of the inputs into the shape (batch_size, n_heads, seq_len, d_head), and somehow handle an output of the shape (batch_size, n_heads, seq_len, d_head), then we could use MHA code something like this: # Q and K are (batch_size, n_heads, len_sequence, d_head) # We need to convert K to (batch_size, n_heads, d_head, len_sequence) # and then we get omega (batch_size, n_heads, len_sequence, len_sequence) omega = Q @ jnp.transpose(K, axes=(0, 1, 3, 2)) omega /= jnp.sqrt(self.d_head) causal_mask = jnp.ones_like(omega, dtype=bool) # tril treats all but the last two axes as batches so we're OK here. causal_mask = jnp.tril(causal_mask) causal_omega = jnp.where(causal_mask, omega, -jnp.inf) # last axis is still OK. attention_weights = jax.nn.softmax(causal_omega, axis=-1) # attention_weights is (batch_size, n_heads, len_sequence, len_sequence) # V is (batch_size, n_heads, len_sequence, d_head) # So this will come out as (batch_size, n_heads, len_sequence, d_head) weighted = attention_weights @ V The next question was, how do we get our inputs into that shape? We could run them all through separate per-head weights -- that is, have an array with one per head, like W_q[0], W_k[0] and W_v[0] for the first one. But that, again, felt like it would be failing to take advantage of the GPU properly. The solution was to think of how matrix multiplications work. If you multiply two matrices, X·Y, the value in the result, in row r, and column c, is the dot product of row r in X and column c in Y. So, imagine if you wanted to multiply X by n different versions of Y, let's call them Y0, Y1, and so on up to Yn. If you imagine a new matrix, Yall, which is basically all the Yxs stacked side-by-side, then the dot-product understanding of multiplication makes it pretty clear that if you did X·Yall, you would get the results of all of those separate multiplications, also stacked side-by-side. I'll call that kind of matrix a "striped" one, for want of a better word. Now, when we project our inputs into the embedding spaces used for attention, we have code like this: Q = self.W_q(xs) We've initialised the weights, W_k in this case, as an nnx.Linear, so what is happening under the hood here is basically: Q=xs·Wq That is, it is just a matrix multiplication. 2 So if we imagine that W_q is one of those "striped" matrices, holding all of the separate matrices to do the projections for all of the heads in a single one shaped (d_emb, n_heads * d_head), then we could stick with the current code -- the Q = self.W_q(xs) Our input xs would be shaped (batch_size, seq_len, d_emb), so the result would be (batch_size, seq_len, n_heads * d_head), and would have the projections for each head in the same vertical stripes as the separate heads' projection weights. Now, like PyTorch, JAX allows you to reshape arrays. You can take one axis of length (say) m×n, and split it into two of lengths m and n respectively -- or, conversely, you can combine two axes of length m and n to one of m×n. If our data had the shape (batch_size, seq_len, n_heads * d_head), we could reshape it like this: Q = self.W_q(xs).reshape((batch_size, seq_len, n_heads, d_head)) ...and that would split things up. So we'd have Q shaped as (batch_size, seq_len, n_heads, d_head). That's almost what we wanted! We needed (batch_size, n_heads, seq_len, d_head), and a simple transpose could sort that out: Q = jnp.transpose( self.W_q(xs).reshape( (batch_size, seq_len, n_heads, self.d_head) ), (0, 2, 1, 3) ) Likewise for K and V, and that was our inputs sorted. Moving on to the output; it came from this: weighted = attention_weights @ V ...and as we worked out above, it was shaped (batch_size, n_heads, seq_len, d_head). I remembered that we wanted to run that through a single linear layer to combine all of the different heads' outputs into one. It felt like the best way to do that would be to get it back into a "striped" layout: (batch_size, seq_len, n_heads * d_head). This would be something like the inverse of the input-wrangling. That would need a reshape, but before I could do that, I'd need to get the axes that needed to be merged next to each other. If the input to the linear layer was going to be (batch_size, seq_len, n_heads * d_head), we'd need to convert it from (batch_size, n_heads, seq_len, d_head) to (batch_size, seq_len, n_heads, d_head) first: jnp.transpose(weighted, (0, 2, 1, 3)) ... and then we could just reshape it to batch_size, len_sequence, n_heads * d_head: striped_output = jnp.transpose( weighted, (0, 2, 1, 3) ).reshape( batch_size, len_sequence, self.n_heads * self.d_head ) Finally, we could run it through a linear layer, with in_features set to n_heads * d_head, and out_features set to d_emb. I put that all together, and decided to throw something extra into the mix. I remembered that Raschka's code had various checks to make sure that d_head * n_heads == d_emb, which seemed a little artificial -- I'd read that this was true of GPT-2, but wasn't a necessary restriction for GPT-style models, which makes sense. There's no obvious reason per se why the heads' embedding dimensions should sum up to the higher-level embedding dimensions. So I decided initially to just pass in d_head and n_heads to the constructor. In my training script I could force them to match the GPT-2 model, but if I wanted to use the code later for something different, I could vary them. Then I remembered that although the dimensionality of the embedding spaces for the query and the key vectors have to match (because otherwise you can't multiply them to work out attention scores with Ω=QKT), the value vector's dimensionality can in theory be different. So I decided to break d_head into two separate d_qk and d_v parameters. The result was this: class MultiHeadAttention(nnx.Module): def __init__(self, d_emb, n_heads, d_qk, d_v, qkv_bias, rngs): self.n_heads = n_heads self.d_qk = d_qk self.d_v = d_v self.W_q = nnx.Linear(d_emb, self.d_qk * n_heads, use_bias=qkv_bias, rngs=rngs) self.W_k = nnx.Linear(d_emb, self.d_qk * n_heads, use_bias=qkv_bias, rngs=rngs) self.W_v = nnx.Linear(d_emb, self.d_v * n_heads, use_bias=qkv_bias, rngs=rngs) self.output_projection = nnx.Linear(self.d_v * n_heads, d_emb, use_bias=False, rngs=rngs) def __call__(self, xs): batch_size, len_sequence, d_emb = xs.shape # For each of the below: # * The initial linear layer projects them to # (batch_size, len_sequence, d_X * n_heads) # where X is qk or v as appropriate. # * The reshape makes them (batch_size, len_sequence, n_heads, d_X) # * The transpose makes them (batch_size, n_heads, len_sequence, d_X) Q = jnp.transpose( self.W_q(xs).reshape( (batch_size, len_sequence, self.n_heads, self.d_qk) ), (0, 2, 1, 3) ) K = jnp.transpose( self.W_k(xs).reshape( (batch_size, len_sequence, self.n_heads, self.d_qk) ), (0, 2, 1, 3) ) V = jnp.transpose( self.W_v(xs).reshape( (batch_size, len_sequence, self.n_heads, self.d_v) ), (0, 2, 1, 3) ) # Q and K are (batch_size, n_heads, len_sequence, d_qk) per above # We need to convert K to (batch_size, n_heads, d_qk, len_sequence) # and then we get omega (batch_size, n_heads, len_sequence, len_sequence) omega = Q @ jnp.transpose(K, axes=(0, 1, 3, 2)) omega /= jnp.sqrt(self.d_qk) causal_mask = jnp.ones_like(omega, dtype=bool) # tril treats all but the last two axes as batches so we're OK here. causal_mask = jnp.tril(causal_mask) causal_omega = jnp.where(causal_mask, omega, -jnp.inf) # last axis is still OK. attention_weights = jax.nn.softmax(causal_omega, axis=-1) # attention_weights is (batch_size, n_heads, len_sequence, len_sequence) # V is (batch_size, n_heads, len_sequence, d_v) # So this will come out as (batch_size, n_heads, len_sequence, d_v) weighted = attention_weights @ V # Transpose to (batch_size, len_sequence, n_heads, d_v), # then reshape to (batch_size, len_sequence, n_heads * d_v) striped_output = jnp.transpose( weighted, (0, 2, 1, 3) ).reshape( batch_size, len_sequence, self.n_heads * self.d_v ) # Final linear layer to combine return self.output_projection(striped_output) Unusually for a case where I went off the reservation like this, the whole thing with the embedding space dimensionality didn't cause any problems at all! But there was one small bug in this code, which I didn't discover until later -- we'll come to it by the end of the post. At this point, I did another of my short training runs, and: 2026-06-23 17:51:32.094308 Tokens seen: 92,209,152 2026-06-23 17:51:32.094311 Throughput: 85,682 tokens/second 2026-06-23 17:51:32.094321 Final train loss: 5.358 2026-06-23 17:51:32.094323 Done ...with a loss chart that looked like this: The training loss at the 937th global step was 5.336, only a tiny bit better than the 5.354 with single-head attention. That was quite possibly within the noise. Even though (due to the d_head * n_heads == d_emb restriction I was enforcing in my training script) the W_q, W_k, and W_v arrays were the same size, I was creating that output_projection, which would consume randomness and make things vary. If I were doing a proper scientific experiment to see if a single layer of MHA beat a single layer of single-head attention, I think I would have run both for more steps to see if the difference became more pronounced later. But for the purposes of this post, I decided to move on. My checklist now looked like this: Convert token IDs into embeddings. ✔ ️done Add on position embeddings. ✔ ️done Run these embeddings through multiple successive Transformers blocks. part-done -- one layer only Layer normalisation ✔ ️done Project them back from embedding space to vocab space. ✔ ️done Inside the Transformers blocks, we: Take a copy of the input sequence of embeddings ✔ ️done Layer normalisation Run multi-head attention ✔ ️done Add the copy back in so that the version that came out of MHA is something more like an "annotation" of the original ✔ ️done Take a second copy of that one Layer normalisation again Run it through a simple neural network Add the results of that back in. Adding that simple neural network -- the FFN -- seemed like a good next step. Adding the FFN to the Transformers block The feed forward network is simple; you take the output of the MHA block, run it through a biased linear layer to expand it from d_emb to 4 * d_emb, then run it through the GELU activation function, then shrink it back down to d_emb with another linear layer. I didn't really see any value in writing my own implementation of GELU, given that even in the book we were just given code for an approximation to type in. So, using jax.nn.gelu, I wrote this: class TransformersLayer(nnx.Module): def __init__(self, d_emb, n_heads, d_qk, d_v, qkv_bias, rngs): self.attention = MultiHeadAttention(d_emb, n_heads, d_qk, d_v, qkv_bias, rngs) self.ffn = nnx.Sequential( nnx.Linear( in_features=d_emb, out_features=d_emb * 4, use_bias=True, rngs=rngs ), jax.nn.gelu, nnx.Linear( in_features=d_emb * 4, out_features=d_emb, use_bias=True, rngs=rngs ), ) def __call__(self, xs): shortcut = xs att = self.attention(xs) post_attention = shortcut + att fed_forward = self.ffn(post_attention) return fed_forward + post_attention Note that I added in a shortcut connection around the FFN as well, so that it didn't overwrite what was there, but only "added on its notes". I kicked that off, and it ran for ten minutes or so, but then OOMed: 2026-06-23 18:20:01.376377 Saving checkpoint 51%|██████████████████████████████████████████████████████▊ | 481/938 [10:20<11:47, 1.55s/it, loss=5.758, tps=76,185]W0623 18:20:12.602631 2860192 bfc_allocator.cc:514] Allocator (GPU_0_bfc) ran out of memory trying to allocate 2.93GiB (rounded to 3149744640)requested by op If the cause is memory fragmentation maybe the environment variable 'TF_GPU_ALLOCATOR=cuda_malloc_async' will improve the situation. Adding TF_GPU_ALLOCATOR=cuda_malloc_async didn't help. I spent some time trying to dig into what might be causing it, but eventually noticed something interesting: in nvtop, the VRAM usage was consistently 75% throughout. Now I knew that JAX pre-allocates 75% of VRAM when it starts up, but I'd been assuming that it would try to grab more if it needed it. It turned out I was wrong with that assumption -- it grabs 75%, but that's all you ever get! The solution turned out to be the XLA_PYTHON_CLIENT_MEM_FRACTION environment variable. If you set that to, say, 0.90, then JAX will pre-allocate 90% of the VRAM, and you can use all of that. (You can also make it allocate as-needed with XLA_PYTHON_CLIENT_PREALLOCATE=false, and there are various other settings you can control with other environment variables on that linked page). Anyway, setting it to 0.90 to grab 90% of VRAM worked, and I was able to get a successful run: 2026-06-24 00:29:34.864880 Tokens seen: 92,209,152 2026-06-24 00:29:34.864882 Throughput: 77,596 tokens/second 2026-06-24 00:29:34.864900 Final train loss: 5.341 2026-06-24 00:29:34.864902 Done The loss chart was this: ...and the training loss at global step 937 was 5.295, compared to the 5.336 from MHA alone. Another tiny improvement, another one that could have been in the noise. Again, if I were doing a proper experiment, I'd do a longer run, but for now, I decided to move on. The checklist looked like this: Convert token IDs into embeddings. ✔ ️done Add on position embeddings. ✔ ️done Run these embeddings through multiple successive Transformers blocks. part-done -- one layer only Layer normalisation ✔ ️done Project them back from embedding space to vocab space. ✔ ️done Inside the Transformers blocks, we: Take a copy of the input sequence of embeddings ✔ ️done Layer normalisation Run multi-head attention ✔ ️done Add the copy back in so that the version that came out of MHA is something more like an "annotation" of the original ✔ ️done Take a second copy of that one ✔ ️done Layer normalisation again Run it through a simple neural network ✔ ️done Add the results of that back in. ✔ ️done Now, my gut instinct was that the layer normalisation inside the Transformers blocks was of most value as a way of stabilising training over deep networks. And with one layer, it didn't seem like the right time to add it. Instead, I decided to add on multiple layers. Multiple layers For GPT-2 small, you have 12 layers. That was already being passed in to my GPTModel's __init__ method as n_layers, so I just replaced this: self.transformers_layer = TransformersLayer( d_emb, n_heads, d_qk, d_v, qkv_bias, rngs ) ...with this: self.transformers_layers = nnx.Sequential( *( TransformersLayer( d_emb, n_heads, d_qk, d_v, qkv_bias, rngs ) for _ in range(n_layers) ) ) ...and then just renamed it where it was called; this: transformed = self.transformers_layer(input_embeddings) ...became this: transformed = self.transformers_layers(input_embeddings) I kicked it off, and it completed! However, the loss chart was telling: Ouch. Loss started dropping quite nicely, but then things got out of control and it settled down at a loss that was essentially that of a random model. At step 937, we were at 10.75, so just a hair less than the 10.82 that randomly guessing next tokens would give. Well, LayerNorm is specifically meant to stabilise training, and the checklist looked like this: Convert token IDs into embeddings. ✔ ️done Add on position embeddings. ✔ ️done Run these embeddings through multiple successive Transformers blocks. ✔ ️done Layer normalisation ✔ ️done Project them back from embedding space to vocab space. ✔ ️done Inside the Transformers blocks, we: Take a copy of the input sequence of embeddings ✔ ️done Layer normalisation Run multi-head attention ✔ ️done Add the copy back in so that the version that came out of MHA is something more like an "annotation" of the original ✔ ️done Take a second copy of that one ✔ ️done Layer normalisation again Run it through a simple neural network ✔ ️done Add the results of that back in. ✔ ️done ...and the only remaining step was that LayerNorm in the Transformers blocks, so it was time to add it in! Adding LayerNorm As per the checklist, we do the LayerNorm after we've taken our copy for the shortcut connection, just before MHA, and then likewise after the second shortcut copy, before the FFN. As I understand it, this was a GPT-2 innovation -- previously, people had done normalisation after those steps, but this pre-norm setup turned out to work better. The code changes were simple. I added two LayerNorm modules to the TransformersLayer class, and then called them in the appropriate places (taking the opportunity to tidy up the variable naming in the forward pass while I was there): class TransformersLayer(nnx.Module): def __init__(self, d_emb, n_heads, d_qk, d_v, qkv_bias, rngs): self.attention_norm = LayerNorm(d_emb) self.attention = MultiHeadAttention(d_emb, n_heads, d_qk, d_v, qkv_bias, rngs) self.ffn_norm = LayerNorm(d_emb) self.ffn = nnx.Sequential( ... ) def __call__(self, xs): shortcut = xs xs = self.attention_norm(xs) xs = self.attention(xs) xs = xs + shortcut shortcut = xs xs = self.ffn_norm(xs) xs = self.ffn(xs) return xs + shortcut I kicked it off and ran it, and got these results: 2026-06-24 03:07:51.128966 Tokens seen: 92,209,152 2026-06-24 03:07:51.128969 Throughput: 23,399 tokens/second 2026-06-24 03:07:51.128979 Final train loss: 5.359 2026-06-24 03:07:51.128981 Done That certainly looked much healthier! However, when I looked at the loss at step 937, it was 5.311 -- a tiny bit higher than the single-layer MHA example, which got 5.295. I'd been willing to play a bit fast and loose with this loss number and allow myself to accept a win when the loss went down a tiny bit, even if it was such a small amount that it could have been within the noise. But increasing loss -- even if it could also be within the noise -- was a step too far. I decided that in this specific case, I'd be strict and test the hypothesis that longer training runs would demonstrate an improvement between one single layer without pre-norm, and multiple layers with pre-norm. I had to remember that these training runs would not be comparable with the earlier ones. In the training script, I had a learning rate schedule like this: That straight-line warmup period and the following cosine decay were 5% and 95% of the training run respectively, which meant that (for example) global step 937 of the short runs we had been doing would be at a completely different point in the schedule than the same step would in these longer runs. However, they would be comparable to each other, and that was what mattered. After some humming and hawing, I decided that a full Chinchilla-optimal (for the full model) training run over 3,260,190,720 tokens, rounded up to fit into a round number of global steps, would be a nice experiment. I expected it to run comfortably overnight for the single-layer run, and take a bit less than two days for the multi-layer one. So I kicked off the first. Just over 11 hours later: 2026-06-24 16:39:52.178045 Tokens seen: 3,260,252,160 2026-06-24 16:39:52.178049 Throughput: 81,733 tokens/second 2026-06-24 16:39:52.178058 Final train loss: 4.324 2026-06-24 16:39:52.178061 Done Here's the loss chart: The last checkpointing period in that run ended at global step 33,164, and the training loss then was 4.165 -- indeed, it had been at around 4.17 for quite some time, though the trend still seemed to be a tiny bit downward. So then I kicked off a run of the full version -- multiple layers, with pre-norm in the Transformers blocks. Just over 37 hours later: 2026-06-26 06:55:39.956609 Tokens seen: 3,260,252,160 2026-06-26 06:55:39.956614 Throughput: 24,151 tokens/second 2026-06-26 06:55:39.956625 Final train loss: 3.637 2026-06-26 06:55:39.956629 Done The "Final train loss" line at the end said it all, really! But here's the loss chart: ...and the loss at step 33,164 was 3.399. Definitely quite an improvement over the 4.165 that a single layer got. Again, at some point I might do the equivalent tests for the earlier results where improvements appear to be pretty much in the noise. It would be good to be sure that the changes really did have the impact I think they did. But for now: our checklist was looking like this: Convert token IDs into embeddings. ✔ ️done Add on position embeddings. ✔ ️done Run these embeddings through multiple successive Transformers blocks. ✔ ️done Layer normalisation ✔ ️done Project them back from embedding space to vocab space. ✔ ️done Inside the Transformers blocks, we: Take a copy of the input sequence of embeddings ✔ ️done Layer normalisation ✔ ️done Run multi-head attention ✔ ️done Add the copy back in so that the version that came out of MHA is something more like an "annotation" of the original ✔ ️done Take a second copy of that one ✔ ️done Layer normalisation again ✔ ️done Run it through a simple neural network ✔ ️done Add the results of that back in. ✔ ️done Everything was checked off. So was this journey over? Well, there was one thing that the original PyTorch code had that my new code didn't: dropout. Dropout I'd found in my lengthy interventions experiments that dropout seemed to make models worse. It was, I felt, a smart idea back in the days when people had little data and did multiple epochs, each sweeping over everything, but it made less sense nowadays with single-epoch training runs over very large datasets. (Though I do have some intuitive ideas about why it could still help.) Still, it would be good to show that it harmed loss for this model as well. Checking my notes, I found that there were four places where dropout was applied: Once in the main body, just after we've worked out the embeddings. Twice in the transformers block: once after attention (but before the shortcut is mixed back in), and once after the FFN (ditto) Inside multi-head attention, on the attention weights (which surprised me). The changes are tiny and rather dotted around the code, so rather than showing you isolated bits of code, if you'd like to see it you can take a look at the code at this point and search for "dropout". When I started running that, I got an error when saving the first checkpoint: TypeError: JAX array with PRNGKey dtype cannot be converted to a NumPy array. Use jax.random.key_data(arr) if you wish to extract the underlying integer array. This was happening deep inside the bowels of Safetensors, but it made a lot of sense. The nnx.Dropout object needs to keep track of the state of the random number generator, and that meant that the to_flat_state function that I was using might return a structure that had something that contained that state, and was not compatible with Safetensors. I decided that I'd cheat a little bit here. If I skipped the dropout layers when I saved my checkpoints, like this: for tuple_key, array in flat_state: key = ".".join(str(key) for key in tuple_key) if "dropout" not in key: simple_dict[key] = array ...then I'd be able to save them. This would have a problem -- if I restarted from a checkpoint, the dropout pattern after the restart would mirror the dropout pattern from the start of the training run, because the random seed it started with would not have come from the checkpoint, but just the initialisation code. I felt that this would not have a serious impact, though, and given that I'd not had to restart from checkpoints so far, I (wrongly, as it turned out) decided it wouldn't matter. I kicked off the run, and... after four hours, it OOMed. I cursed, decided that I'd nurse this run through anyway (despite my dropout checkpointing concerns), and kicked it off again. Three hours later, it OOMed again. I happened to be away from home at the time, logging in to my machine remotely (thanks, Tailscale!), and on looking at nvtop, I realised that the X window system on my machine was using a gig or so of VRAM. I was running the training run in a tmux session, which meant that I could kill X and not lose state, so I did that, and adjusted the XLA_PYTHON_CLIENT_MEM_FRACTION environment variable I was using -- it had been 0.90, so I bumped it up to 0.95. I kicked it off again, and... 2026-06-28 22:06:47.669676 Tokens seen: 2,640,052,224 2026-06-28 22:06:47.669683 Throughput: 23,019 tokens/second 2026-06-28 22:06:47.669691 Final train loss: 3.776 2026-06-28 22:06:47.669694 Done Note that the tokens seen only relates to the period since the restart, which is why it was lower. One more loss chart: ...and the training loss at step 33,164 was 3.524, higher enough than the 3.399 I got without dropout that I was comfortable that it wasn't in the noise. That was very reassuring. Once again, if this was a proper scientific experiment I'd fix the issue with saving dropout, and run it completely from scratch -- or, at least, run it all the way through from scratch without restarts, even if I had to try several times to get it done. But I don't think that "replaying" dropout would make the loss any worse. And for this experiment, I felt this was enough. So: checklist complete. GPT-2 model coded up. It was time for some evals! Evals, first try -- and fixing an MHA bug I wanted to evaluate these models against the ones I got using the old PyTorch code: specifically, the last local training run that used exactly the same training hyperparameters, and only differed in that it was trained using AMP -- 32-bit floats in general, but using 16-bit where the framework thought it would not be harmful. In order to do exactly the same evals, I decided it would be easiest to write a conversion script to take the Safetensors files written to my JAX checkpoints, and write out new files that were compatible with the PyTorch model code -- then I'd be able to use the original PyTorch eval code. I put something together, converted my last two models -- the full runs with and without dropout -- and tried to load them up. Unfortunately there was an error: RuntimeError: Error(s) in loading state_dict for GPTModel: Missing key(s) in state_dict: "trf_blocks.0.att.out_proj.bias", "trf_blocks.1.att.out_proj.bias", "trf_blocks.10.att.out_proj.bias", "trf_blocks.11.att.out_proj.bias", "trf_blocks.2.att.out_proj.bias", "trf_blocks.3.att.out_proj.bias", "trf_blocks.4.att.out_proj.bias", "trf_blocks.5.att.out_proj.bias", "trf_blocks.6.att.out_proj.bias", "trf_blocks.7.att.out_proj.bias", "trf_blocks.8.att.out_proj.bias", "trf_blocks.9.att.out_proj.bias" You might remember that back when I went through multi-head attention, I mentioned that I'd made a mistake. Somehow, I'd misremembered, and thought that the output projection -- the one that mixes together all of the different heads' outputs -- was a linear layer without bias, despite my original notes being perfectly clear that it did have bias. The good news was that if I disabled bias in the PyTorch code, I could load the safetensors files that I had. So the two models I'd trained so far were not useless, and could actually work as a kind of natural experiment into the benefits of having that bias there. But anyway, in order to do things properly, I was going to need to fix the bug and train yet another model. Adding bias to the MHA output projections The fix was simple, I just replaced this (in MultiHeadAttention): self.output_projection = nnx.Linear(self.d_v * n_heads, d_emb, use_bias=False, rngs=rngs) ...with this: self.output_projection = nnx.Linear(self.d_v * n_heads, d_emb, use_bias=True, rngs=rngs) Then it was time to kick off yet another training run. After another 37 hours: 2026-07-05 10:09:22.147819 Tokens seen: 3,260,252,160 2026-07-05 10:09:22.147823 Throughput: 24,072 tokens/second 2026-07-05 10:09:22.147832 Final train loss: 3.650 2026-07-05 10:09:22.147834 Done ...with this loss chart: ...and the training loss at step 33,164 was 3.398 -- almost exactly the same as the 3.399 that I got in the no-dropout training run without MHA bias above! Well, now it really was time for the evals. Evals, take two I updated my conversion script to handle the bias on the MHA output projections, and used it to convert the three models -- the un-biased ones, with and without dropout, and the biased one, without -- to the PyTorch format, then ran the loss test that I had been using to compare the old models on each. Here are the results, compared to the previous models, and OpenAI's: Test loss OpenAI weights: medium 3.231442 JAX, with MHA bias, no dropout 3.418784 JAX, no MHA bias, no dropout 3.420089 JAX, no MHA bias, with dropout 3.476802 OpenAI weights: small 3.499677 1xrtx3090-stacked-interventions 3.538161 8xa100m40-stacked-interventions-1 3.577761 Cloud FineWeb, 8x A100 40 GiB 3.673623 1xrtx3090-baseline 3.683835 8xa100m40-baseline 3.691526 Cloud FineWeb, 8x H100 80 GiB 3.724507 Cloud FineWeb, 8x A100 80 GiB 3.729900 Cloud FineWeb, 8x B200 160 GiB 3.771478 Local FineWeb train 3.943522 Local FineWeb-Edu extended train 4.134991 Local FineWeb-Edu train 4.166892 That was a pretty amazing result -- I'd clearly proven that JAX trains much better models than PyTorch! 3.5% better in the best case. Well, OK, no. My guess is that the difference was probably something like better luck with the initial weights on the JAX side, plus the improvement from not using AMP. Anyway, the important thing was that the JAX models were in the same kind of loss range as the PyTorch ones -- and while a 3.5% improvement in loss was more variation than I'd been expecting, it was definitely the right ballpark. Now, one thing I had found in the past was that the OpenAI weights -- and some of my own models, like the Fineweb-Edu ones -- were consistently better at an instruction fine-tuning test than their test loss scores would indicate. Would that hold here? The IFT eval code fine-tuned each model on the Alpaca dataset until validation loss started rising, then used the model prior to the start of the rise to generate responses for a test set. These were saved, and then run past an OpenAI model so that they could be compared with each other: You are judging the comparative capabilities of a number of different LLM models. They have been trained to follow instructions. The input was this: ` {input} ` An example correct output is this: ` {correct_output} ` Please produce a score of between 0 and 100 for each model, and respond with a JSON structure like this (note that the number of models may differ from this example): ` { "Model 1": {"score": XXX, "comments": "optional comments"}, "Model 2": {"score": YYY, "comments": "optional comments"}, "Model 3": {"score": ZZZ, "comments": "optional comments"} } ` ...where the XXX, YYY and ZZZ are the scores for the respective models. You can optionally add the "comments" field if you want to explain your reasoning. Here are the models' responses: # Model 1 {model 1 response} # Model 2 {model 2 response} # Model 3 {model 3 response} ...with the model order randomly changed for each query to avoid any position bias. The methodology seemed solid, but I was uncertain about the "train until loss starts rising", as it meant that different models had wildly different amounts of fine-tuning -- between two and seven epochs. On the one hand it felt "unfair" to certain models that they'd get less training than others. On the other hand, if the less-trained models had been trained past the point where their validation loss started rising, then assuming that loss would continue to rise, further training would actually be a disadvantage rather than an advantage. I decided to stick with the original plan, and train until validation loss started rising. I did, however, switch the judge model from the GPT 5.4 that I used in my last IFT test to GPT 5.5. Here are the results: Test loss IFT epochs IFT score IFT rank OpenAI weights: medium 3.231442 2 41.62 1 JAX, with MHA bias, no dropout 3.418784 4 19.25 4 JAX, no MHA bias, no dropout 3.420089 3 14.66 11 JAX, no MHA bias, with dropout 3.476802 4 12.94 15 OpenAI weights: small 3.499677 2 26.73 2 1xrtx3090-stacked-interventions 3.538161 4 17.79 6 8xa100m40-stacked-interventions-1 3.577761 4 10.29 16 Cloud FineWeb, 8x A100 40 GiB 3.673623 7 20.71 3 1xrtx3090-baseline 3.683835 6 15.11 9 8xa100m40-baseline 3.691526 4 14.74 10 Cloud FineWeb, 8x H100 80 GiB 3.724507 4 13.25 14 Cloud FineWeb, 8x A100 80 GiB 3.729900 4 14.50 12 Cloud FineWeb, 8x B200 160 GiB 3.771478 4 16.03 8 Local FineWeb train 3.943522 7 13.73 13 Local FineWeb-Edu extended train 4.134991 7 16.70 7 Local FineWeb-Edu train 4.166892 7 18.68 5 More interesting datapoints! As before, you can see that low loss is not particularly well-correlated with a high score on this instruction fine-tuning test. The OpenAI weights continue to lead the pack, and while one of our new JAX models did quite well, it's still beaten by the Cloud FineWeb, 8x A100 40 GiB model. But what was important here, just as with the loss, was that the new JAX models landed in the same ballpark as the PyTorch ones. They did, and so I could be confident that they were doing essentially the same thing. And that meant that, after 18 months, I had reached the end of my LLM from scratch journey. Conclusion It's been a long trek. I started reading "Build a Large Language Model (from Scratch)" on 22 December 2024. I was planning to breeze through over the Christmas break, but somehow it morphed into being a curriculum onto which I could hang projects to learn the fundamentals of LLMs, beyond what was in the book. In May 2025, I had my first real conceptual breakthrough when I realised that attention heads are (individually) dumb, and as I continued, the second big one came later on in the same month, when the concept of embeddings as being projections between vocab space and embedding space (and the converse projection in the other direction that happens in the LLM's output head) became clear. In August I had the first moment where I felt that the standard teaching approach to LLMs might not be the full story; shortcut connections are normally explained as a way to fix vanishing gradients, while I felt that a better way to see them was a way to allow attention and the FFN to "annotate" the existing information, similarly to how Jewish scholars have annotated the original text of the Talmud. (The results in this post seem to point in that direction, given how even a single layer of attention was massively helped by adding them.) By early December, I had essentially finished the book, and felt I wanted to try to train my first base model from scratch on my RTX 3090. It worked, and wasn't far off the quality of the original GPT-2 small. I was really surprised that I could do that with consumer hardware, and became interested (perhaps obsessively so) with whether I could match OpenAI's weights. In January 2026, I trained a model using DDP on Lambda Labs, and then spent the following months training model after model, trying to work out which interventions -- learning rate scheduling, gradient clipping, etc -- would improve the loss. I wrapped that up in late April, with the interesting finding that although I'd been able to get the test loss pretty low, that didn't seem to map cleanly to performance in my instruction fine-tuning tests. In other words, Loss Number Goes Down is an interesting technical game to play, but doesn't cleanly map to real-world performance. The final step was this post, and the previous one -- could I, using my notes, implement GPT-2 completely from scratch in JAX without referencing the book? And as you've read, the answer was a definite yes! Of course, as with any long-running project, there are some loose ends -- from this post alone, there's the interesting fact that JAX trained faster than PyTorch (perhaps torch.compile could close the gap?) and had a larger possible batch size for full-fat 32-bit. And the fact that fixing the multi-head attention bias bug didn't seem to help with the loss much was interesting too. But those are really details, and there's so much beyond them to learn. Longer-context LLMs: position embedding improvements like RoPE, efficiency tricks like flash attention and attention variants like DSA. Mixture of experts models. How do optimisers really work? (Do they work?) And plenty more. So it's time to draw a line under this series, and start thinking about what comes next. It's been a blast; if you've been reading along, I hope it's been as useful (and fun) to read as it was to write. And as always, comments, questions and corrections very welcome below. On looking back at Raschka's code, after having worked through all of this, there's a slight difference. I do this: normalized = (xs - means) / (stds + 1e-5) ...whereas he does this: norm_x = (x - mean) / torch.sqrt(var + self.eps) Now, the standard deviation is the square root of the variance, so if you ignore the small numbers -- 1e-5 in my case, and self.eps in his -- the calculations are the same. But there is a difference once those are taken account of. I don't think it's large enough to have any serious effect in these runs, though. ↩ In PyTorch, linear layers are stored as the transpose of the matrix that would allow you to do that, so it would be: Q=xs×WqT Also, note that for simplicity (heh) I'm disregarding bias in this discussion. ↩

8th Jul 2026 • 2 votes
Using Safetensors with Flax

I'm porting my PyTorch LLM code to JAX, using Flax as the neural network layer. For various reasons I wanted to use Safetensors to store checkpoints of the model. It took a little while to get it working; here's the trick I learned. If you look at the Safetensors docs, you'll see that it doesn't mention a JAX implementation -- indeed, searching for "safetensors jax" at the time I'm writing this gives you a link to this GitHub repo by Alvaro Bartolome -- which was last updated in 2023. However, if you look more closely at the docs, they do have a link to the Flax API. I feel this is somewhat misnamed, as it is actually a JAX API. There's no reference (again, as of the time of writing) to Flax in the source -- it's all just JAX code. And in fact Bartolome's library uses it under the hood. There is one problem, though. The API works with simple single-level dictionaries, with strings mapping directly to JAX arrays. For example, the save_file function has this signature: def save_file( tensors: Dict[str, Array], filename: Union[str, os.PathLike], metadata: Optional[Dict[str, str]] = None, ) -> None This can cause problems if you're not careful. If you look at the Flax documentation on checkpointing, it suggests that you use Orbax 1, which has its own API and file format, but then goes on to say: When interacting with checkpoint libraries (like Orbax), you may prefer to work with Python built-in container types. In this case, you can use the nnx.State.to_pure_dict and nnx.State.replace_by_pure_dict API to convert an nnx.State to and from pure nested dictionaries. I initially put two and two together -- that and the dictionary-based API for Safetensors -- and got five, and tried feeding one of those "pure" dicts into Safetensors. I got a very confusing error: SafetensorError: dtype object is not covered It's worth digging in to why that happens. The problem is that although Safetensors is expecting a dict of strings mapping to tensors, it doesn't check that that is what it actually gets. And while the dictionaries from nnx.State.to_pure_dict are "pure", they are also nested (as the docs say!). Even for the simple model I was working with, I got a structure like this: { 'output_head': { 'kernel': Array([...], dtype=float32) }, 'token_embedding': { 'embedding': Array([...], dtype=float32) } } So, we had strings mapping to dicts, and those dicts mapped from strings to the JAX arrays. More complex models would have had deeper dict structures. Now, internally inside Safetensors, the Flax/JAX API is a simple wrapper. It iterates over the keys in the dictionary it's been provided with, and tries to convert their respective values into NumPy arrays. It does that by passing them into NumPy's asarray function, which accepts things like lists, tuples, and NumPy arrays, and converts them into arrays. JAX's own Array class exposes an interface that it recognises, so they're converted without trouble. Once it's done that, it passes the result to a lower-level Rust implementation that actually converts everything to Safetensors format. But because Safetensors didn't check types, in my case it was iterating over the top level of the dict, trying to convert the values to NumPy arrays, and got something like this: { 'output_head': numpy.array({'kernel': Array([...], dtype=float32)}, dtype=object), 'token_embedding': numpy.array({'embedding': Array([...], dtype=float32)}, dtype=object) } That is -- because it assumed that the values in the top-level dict were JAX Arrays, it blindly tried to convert them to NumPy arrays. But they were dicts (that happened to map from strings to arrays) -- and if you ask asarray to create an array based on a random object, it happily does so and wraps that object in a NumPy array, with a dtype of object. When that is then fed into the lower-level Rust code that is trying to write the file, it encounters NumPy arrays that have a dtype it can't handle, object -- hence that error: SafetensorError: dtype object is not covered It all makes sense when you read through the code, but I was a bit perplexed for a while! I think all this might be the reason why Bartolome created his GitHub repo. In the README, he says that: There are no plans from HuggingFace to extend safetensors to support anything more than tensors e.g. FrozenDicts, see their response at huggingface/safetensors/discussions/138. So the motivation to create safejax is to easily provide a way to serialize FrozenDicts using safetensors as the tensor storage format However, you don't need to use that library to serialise simple Flax models. Consider how PyTorch models get serialised to Safetensors; my LLMs have keys with names like out_head.weight, pos_emb.weight, and trf_blocks.0.att.out_proj.weight. They're "flat" dictionaries mapping strings to PyTorch Tensors, similar to what Safetensors wants for these Flax ones, but they use dots to separate different levels, with integers for list items and strings for field names. Looking at the pure-dict structure I had for my model: { 'output_head': { 'kernel': Array([...], dtype=float32) }, 'token_embedding': { 'embedding': Array([...], dtype=float32) } } ...you can see that you could walk the dictionary structure to generate keys like output_head.kernel and token_embedding.embedding. That would be easy enough to code up. But -- as Adithya Dsilva points out on GitHub -- you can get there even faster by using nnx.to_flat_state. That returns a (non-dict) structure like this: FlatState([ (('output_head', 'kernel'), Param( # 786,432 (3.1 MB) value=Array([[ 2.3581974e-02, 3.0957451e-02, -3.5088759e-02, ..., -4.5880198e-02, 5.3717274e-02, -2.6590331e-02], ..., [-9.6302675e-03, -3.3276502e-02, 5.7173111e-02, ..., -7.9063717e-03, 2.0532632e-02, 5.4753982e-02]], dtype=float32) )), (('token_embedding', 'embedding'), Param( # 786,432 (3.1 MB) value=Array([[ 0.00273973, -0.01754938, 0.04656043, ..., -0.04276522, -0.03986642, -0.00781331], ..., [ 0.01421758, -0.0219186 , -0.01701825, ..., -0.00793659, 0.00500103, 0.03839901]], dtype=float32) )) ]) If you iterate over that FlatState, you get tuples where the first element is that tuple of strings, like ('output_head', 'kernel'), and the second is a Param object wrapping the JAX Array. The tuples mirror the dot-separated string format in the PyTorch-style Safetensors files. Param objects also implement an interface that asarray can understand, so you can quickly and easily convert the FlatState to a regular dict for Safetensors: from safetensors.flax import save_file ... model_state = nnx.state(model) flat_state = nnx.to_flat_state(model_state) simple_dict = {} for tuple_key, param in flat_state: key = ".".join(str(key) for key in tuple_key) simple_dict[key] = param save_file(simple_dict, "model.safetensors") (You need to wrap key in a str because if you have a nnx.Sequential in your model, the item in the tuple will get an integer index rather than a string). You can go the other way pretty easily too; given a model, you can load the saved checkpoint into it like this (because from_flat_state accepts raw JAX Arrays in place of explicit Params): from safetensors.flax import load_file ... simple_dict = load_file("model.safetensors") dict_flat_state = {} for key, array in simple_dict.items(): elements = key.split(".") list_key = [] for element in elements: try: list_key.append(int(element)) except ValueError: list_key.append(element) dict_flat_state[tuple(list_key)] = array new_flat_state = nnx.from_flat_state(dict_flat_state) nnx.update(model, new_flat_state) A little more work than I'd ideally like, but given that it can be tucked away in general save_checkpoint/load_checkpoint functions, not too big a deal. Hope that's of use for other people coming across this problem! I'm beginning to feel a bit swamped with all of these libraries with names ending in -ax. It reminds me of the names of the characters in Asterix's village... ↩

4th Jun 2026 • 1 votes

More in AI

Dyson CameraJet

So when I saw Dyson had a $500 toothbrush, I was excited. Finally, advertising that targets me! I love brushing my teeth, and I have more money than I know how to spend. Not because I’m particularly rich, but because most stuff doesn’t really appeal to me. Like if I owned a helicopter it would just be a headache, because like imagine one day I get a call from the hangar saying the hangar is flooding and the water is rising and you need to move your helicopter. I’m thousands of miles away and need a helicopter pilot in the next 30 minutes, a new place to store it, was the maintenance even done will we even be able to take off on short notice and really I just am upset with myself because I made the poor decision to purchase a helicopter, and once I come back to reality I feel relieved that I don’t own a helicopter and this scenario will never happen to me. I do however, by means of my birthday, own a Dyson CameraJet (pictured above). It broke within 30 seconds of the first brushing. None of the LEDs turn on anymore. I spent an hour investigating, finally opening the user removable battery compartment to find the Spearmint Dyson Low-foaming mouth rinse had leaked inside. And by how the toothbrush is designed, it’s clear the entire electronics compartment was flooded with the stuff. Here’s the top comment on Reddit about this toothbrush. Apparently this is happening to everyone, “a potential for water seepage” they say. Dyson wants me to find the receipt and return it through some obtuse process that probably doesn’t work, dude it was a gift I just want my $500 toothbrush to work. They claim they worked on it for 6 years, but it’s clear their QA Process doesn’t include putting any liquid in the device. It clearly should, ideally for all devices but at least for spot checks on some. It’s sad to see this. At comma, we put every comma four in a highly stressful environment for 16 hours, a superset of the state it’s in driving, while testing all peripherals: the camera, IMU, GPS, screen, etc… We have gotten the failure rate super low by doing this, and for the few that do fail it’s usually after a while. There’s no excuse for a mature consumer electronics company to not design a procedure to fully test the functionality of each device before shipping. This shows some serious dysfunction at the company, and they should take this as a wake up call to fix their processes and issue a recall for the toothbrush. Dyson, if you see this post, e-mail me when I can drop by the Dyson store in ifc mall Hong Kong and swap it for a new one. I don’t want a stupid process, I want a real technical explanation of the issue and a working fancy toothbrush.

2 days ago • 1 votes
Why do OpenAI's GPT-2 weights beat mine? Part five: data quality

When I finished learning how to build an LLM from scratch, I was left with a mystery: my own models were not as good as OpenAI's original GPT-2 models, despite being based on the same architecture. My models all had 163M parameters, and followed the design from Sebastian Raschka's book "Build a Large Language Model (from Scratch)". That meant that they were pretty much the same as the setup for the OpenAI GPT-2 "small" instance, except that they did not use weight-tying or bias on the QKV matrices. Weight-tying means that you re-use the initial embedding matrix as the output head at the end, and using it means that GPT-2 small saved quite a few parameters -- it was 124M rather than 163M -- at, at least in my own experiments, a cost in quality; similarly, while I found that QKV bias made a tiny improvement in loss terms, I'd felt it was likely within the noise. But GPT-2 small consistently beat my models on an instruction fine-tuning (IFT) task -- also adapted from Raschka's book. That test fine-tunes the model on a subset of the Alpaca dataset, until validation loss starts rising, and then runs a test set through the resulting model. The responses to the test set questions are stored, and then I run all of the responses from all of the models under test past GPT 5.5 in one go to get an aggregate score; more details here. GPT-2 small always did better than any of my models on this. Additionally, it did surprisingly well on a simpler eval -- one that just measured the cross entropy loss it got on a test set. It scored close to my own best models, and better than many of them. What made this result particularly interesting was that the test set in question was a split of my own training data; my models would not have seen it when training (at least, in theory), but it seems likely that it would be much more similar to their own training data than it was to OpenAI's. I've checked two things while probing this mystery: It seems very likely that the GPT-2 models were overtrained by modern standards; would overtraining my own models get them closer? It turned out that no, it probably didn't help with the IFT eval (though there might have been some signal there). It did help quite a lot with the test loss eval, though. The way I was handling dropout in the IFT test might have been unduly benefiting some models while working against others. I decided to standardise on not using dropout during this eval, as (counter-intuitively for me) it seemed to harm the results of most models, even those that had been pre-trained with dropout. In particular, the OpenAI weights were harmed by using dropout, and making a change that benefited them (along with some of my own models) seemed the most conservative approach to take in investigating this. The next thing I wanted to look into was the training data. The exact dataset that the various GPT-2 models were trained on has never been released; all we know about it is from the paper, where they say: [W]e created a new web scrape which emphasizes document quality. To do this we only scraped web pages which have been curated/filtered by humans. Manually filtering a full web scrape would be exceptionally expensive so as a starting point, we scraped all outbound links from Reddit, a social media platform, which received at least 3 karma. This can be thought of as a heuristic indicator for whether other users found the link interesting, educational, or just funny. They called it "WebText". There is an OpenWebText that tries to replicate it, but although they tried to follow the same procedure as the original, there's no guarantee that it is all that similar. By comparison, I'd normally been training against FineWeb. While this is a general web-scraping dataset, without the "curation" provided by using only stuff that was linked from upvoted Reddit posts, it has been refined to remove any obvious junk. I had felt that it was pretty much equivalent. But what if I were wrong about that? I decided to see if I could get better models by using better data. The starting point Here's a table of all of the models I've been comparing to date. The "Test loss" column shows how well the model in question did on that held-back cross entropy loss evaluation. The "IFT epochs" column shows how many epochs of fine-tuning the model needed before its validation loss started rising, the "IFT score" the score that GPT 5.5 gave the model's responses to the test set of my Alpaca data, and the "IFT rank" the model's rank in terms of that score. The OpenAI small model is in there in bold, and I've also included the OpenAI medium model for comparison purposes. Test loss IFT epochs IFT score IFT rank OpenAI weights: medium 3.231442 2 43.75 1 JAX, overtrained one long epoch 3.324953 3 19.77 4 JAX, overtrained two normal epochs 3.326482 4 19.72 5 JAX, with MHA bias, no dropout 3.418784 4 18.69 6 JAX, no MHA bias, no dropout 3.420089 5 21.46 3 JAX, no MHA bias, with dropout 3.476802 5 13.22 15 OpenAI weights: small 3.499677 2 26.00 2 1xrtx3090-stacked-interventions 3.538161 4 13.77 14 8xa100m40-stacked-interventions-1 3.577761 4 10.76 18 Cloud FineWeb, 8x A100 40 GiB 3.673623 3 17.72 7 1xrtx3090-baseline 3.683835 4 15.74 8 8xa100m40-baseline 3.691526 3 14.19 13 Cloud FineWeb, 8x H100 80 GiB 3.724507 4 14.33 12 Cloud FineWeb, 8x A100 80 GiB 3.729900 3 11.34 17 Cloud FineWeb, 8x B200 160 GiB 3.771478 4 14.67 11 Local FineWeb train 3.943522 5 12.31 16 Local FineWeb-Edu extended train 4.134991 5 15.04 9 Local FineWeb-Edu train 4.166892 5 14.99 10 You can see that the OpenAI small model did pretty well in terms of the test loss, when you consider that it has 39M fewer weights than my models and was being tested against a dataset that differs more from its likely training data than it does from my own models'. Additionally, the specific models that did better than OpenAI's small one were all trained with JAX rather than PyTorch -- my hypothesis for that is that it's a result of the JAX ones getting better initial weights by pure chance. But the big difference was in the IFT score. In the specific run that gave the results in this table, the OpenAI small model got 26.00 -- the closest of my own models was more than 4.5 points lower, at 21.46. This difference was consistent over all of my other test runs. The GPT-2 small model was always ahead of mine. (GPT-2 medium, of course, beat GPT-2 small and all of my models, but given that it is twice the size of mine, that's not a big surprise.) Now, quite some time ago, I had tried looking into data quality as a lever to pull for model performance. At the bottom of the table, with the worst test loss of all models, you can see two models: "Local FineWeb-Edu train" "Local FineWeb-Edu extended train" These two were (as you might guess from the names) trained on the FineWeb-Edu dataset, which includes just the most "educational" data from FineWeb. They scored very badly on the test loss score. Given that the test dataset is from FineWeb, that's not a big surprise -- as I've written previously: If you train a model on Jane Austen and then evaluate against Chuck Tingle, then you're not going to get amazing results. But again, GPT-2 had the same issue, and did perfectly well on the test loss eval. On the other hand, while these FineWeb-Edu models' performance on the IFT eval wasn't stellar -- there are plenty of my other models ahead of them -- they did seem to punch above their weight. Consistently across all of the IFT evals I've done, they have scored higher than many of the others -- despite their poor loss on the test eval. Additionally: they were amongst the first models that I trained, before I'd spent time learning about how to optimise my hyperparameters and training loop. They did not use gradient clipping, they did use dropout, their batch size was just "whatever I could squeeze into the GPU", and I didn't set the learning rate to the right kind of value or schedule it over the course of the training run. So maybe a new training run on FineWeb-Edu plus my training improvements would help? And maybe some other tweaks to the training data would be worth looking into? The plan I decided to see what would happen if I trained some models with better-quality data. Specifically, I would train models with my current optimised loop and hyperparameters on four different datasets: FineWeb-Edu -- essentially the same as "Local FineWeb-Edu train" but with a better training setup. This would test the "more educational -> better" hypothesis. A 50:50 split of FineWeb and FineWeb-Edu. I've read that LLMs can be helped by having a decent amount of lower-quality data in their training loop, as it helps them to generalise. Perhaps having some FineWeb in there in addition to the FineWeb-Edu stuff would improve that test loss score while also helping the IFT test? A "curated" dataset containing 45% of its contents from FineWeb, 45% from FineWeb-Edu, and 10% from the Simple English Wikipedia. The full Wikipedia is huge, and full of obscure facts -- while the Simple English one is small and hopefully richer in useful information on a per-token basis. And conveniently, Answer.ai have made a snapshot of it available on Hugging Face Hub. Might deliberately putting a bunch of encyclopaedic data into the training set make the model better at the IFT eval (which has lots of factual questions in it, like "who wrote Pride and Prejudice")? OpenWebText. Even though I was unsure how well it matched the original WebText, given that it was there, it seemed silly to not try training something on it and see how it matched up. I would train each model on 3.2B tokens of the chosen dataset; that's the Chinchilla-optimal amount for my 163M-parameter models. If there were any interesting results, then I might consider doing overtrained models later on. I decided to be at least vaguely scientific about this, and to pre-register some predictions: The FineWeb-Edu-only model would do pretty badly on the test loss, but better than my older FineWeb-Edu models (90%). It would also punch above its weight on the IFT eval (90%). The 50:50 split: I expected it to do worse on the test eval than my JAX FineWeb-only models (70%), but better than the FineWeb-Edu one (90%). I wasn't sure about how it would do on the IFT eval, but thought it might be somewhere in between the two groups (60%). The curated dataset I had high hopes for in terms of the IFT eval -- let's say 80% chance of it being the best of all of my models. For the test loss eval, I expected it to do about as well as the 50:50 split, maybe a little bit worse (70%). I had no idea how the OpenWebText eval would do! Could be worse, could be better. Here's how things turned out. The FineWeb-Edu model I already had a dataset based on FineWeb-Edu ready to go, from when I trained those two original models. It is just the 10B-token sample of the original dataset at the time I generated it last December, formatted appropriately for my training script (details on the dataset card). I kicked off a training run with my JAX code (which I've been using for the other posts in this series): 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] ...and just less than 40 hours later, I had a model: 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 I converted the saved JAX safetensors file from the last checkpoint into a format that would be compatible with my PyTorch eval code, and ran my smoke test: how would it complete the sentence "Every effort moves you"? Every effort moves you closer to God’s Kingdom, and even closer to Him. As we can see in That was nice and coherent -- if unusually religious! -- so that was promising. I ran the test eval: 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 That was pretty good, putting it at a better test loss than all of the models I had trained without optimised hyperparameters, and worse than all of the ones I had trained on FineWeb with optimised hyperparameters. So that fit in with my prediction that it would be better than the old FineWeb-Edu models; the fact that it was also better than the non-optimised training runs with FineWeb seemed sensible enough that I felt silly for not having predicted that it would have fallen exactly there :-) I decided to leave the IFT eval until the end so that I could check all of the models from these experiments together, so it was time to upload this one to Hugging Face, and move on to the next model. 50:50 FineWeb to FineWeb-Edu I put together a new repo with a script to prepare datasets specifically for my training setup. You provide it with config that specifies some source datasets along with information about how to process them and how to mix them together, and it uploads a new dataset to Hugging Face Hub with the required characteristics. For example, for the 50:50 FineWeb to FineWeb-Edu split, the config looked like this: { "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 } ] } The way the script works is pretty simple: it works out (based on those weights and the tokens_desired) how many tokens it wants from each source dataset, shuffles the items in the sources, then it loops until it has the desired number of tokens or more stored in an output. In the loop, it works out which source is currently most under-represented, grabs an item from it, tokenises it, and adds it to the output. Running it with that 50:50 config seemed to work fine: 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 So we had almost-perfect 50:50 balance between the datasets, and it saved this dataset on Hugging Face. I ran a script to double-check that it looked sane, and it did, so it was time to spin up a training run: 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] That was running on perry, my normal workstation, and I kicked it off in parallel with the "curated" model training run below on poppy my training box, but I'll keep the runs separate for the purposes of this writeup. When this had been running for an hour or so, our power went out. My guess is that having the tumble dryer running, the car charging, the kettle boiling, the electric hob switched on, and two machines doing training runs is a bit too much for our electrics... which might be a problem in the future, especially if (as planned) I make poppy a multi-GPU machine. However, as things stand, I was able to kick it off again after switching the circuit breaker back on, and things held up. Again, about 40 hours later: 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 (Note that the numbers reported at the end of a restarted run like this only include what happened after the restart.) I converted it to PyTorch-compatible tensors, and did the smoke test: Every effort moves you on to other options—in fact, it’s not even worth that effort. Just make Looking good! Time for the loss test: 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 That was almost in keeping with my prediction that it would do worse than the JAX FineWeb-only models, except that it was better than the worst of those, "JAX, no MHA bias, with dropout": it was actually better than I predicted. So, a promising model. Time to upload it to Hugging Face -- and now let's move on to the next one. The "curated" dataset With my dataset-preparation script, this was easy enough to set up: { "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 } ] } Running that worked nicely: 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 One thing that is worth noting in that output is the "6 iterators" for the Simple English Wikipedia. If a source dataset runs out of items while we're building up the results in this script, we start iterating over it again (with a different seed for the shuffle so that the ordering is different). The "6 iterators" means that it needed to do that 6 times -- the original creation of the iterator at the start of the script, and five more. So that means that the Simple English Wikipedia is repeated (oversampled) somewhere between five and six times in the dataset. That's not a bad thing! From what I've read, it's actually quite standard to oversample highly educational content in LLM training datasets. And anyway, the dataset the script generated was 10B tokens, of which we're only using 3.2B for the training run in this post, so it would only appear somewhere between one and two times. The repetition would likely only really cut in if and when we did an overtrained model on the dataset. Anyway, I ran my check against the uploaded dataset -- the first few items were clearly from FineWeb, FineWeb-Edu, and the Simple English Wikipedia. It was time to kick off a training run: 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] Again, this was interrupted by the power outage that hit the 50:50 training run, but I was able to restart from a checkpoint. After another 22 hours, it crashed with an error that I've seen before: 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'] I put it aside as a one-off oddity when I hit it last time, but this time I dug in a bit more. I noted that it had not ever happened on perry, but seemed to be an issue on poppy, and that poppy had an older version of CUDA and the Nvidia drivers -- might that be the cause? I decided to upgrade those before kicking off the next run, but for now just restarted the run from the most recent checkpoint. (Note for anyone who is hitting the same error: it has not occurred since the upgrade, so that's worth trying.) This time it completed OK: 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 Again, these numbers just show what happened after the most recent restart. I copied it over to perry, converted it into a format that was compatible with my PyTorch code, and ran the smoke test: Every effort moves you by the air, for it will make you a better athlete, so your body becomes bigger and stronger Coherent enough -- time for the loss eval: 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 Again, in line with my predictions -- worse than the JAX FineWeb-only models, and indeed than the very best PyTorch one, 1xrtx3090-stacked-interventions, and also worse than the 50:50 split, but better than the FineWeb-Edu one. I uploaded it to Hugging Face, and it was time to move on to what was meant to be the final model for this set of experiments. The OpenWebText run Again, this was a simple enough config to set up: { "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 } ] } ...and the build and upload process worked well (and took much less time -- for some reason, sampling randomly from a single dataset is faster than sampling from two or three): 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 Note that it needed to oversample -- that "2 iterators". OpenWebText is about 40 GiB uncompressed, and so that's about 10B GPT-2 tokens -- presumably just a little bit less. Again, given that I was planning to use just the first 3.2B tokens of the dataset, I didn't feel that it would matter. I ran the check script on the newly-uploaded Hugging Face dataset and all looked well, so that was all set for the training run. I upgraded poppy first with a sudo pacman -Syu to see if that helped with the weird error that I got in the previous run (which, as I said, it looks like it did), then kicked it off: 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] About 31 hours in, it crashed again, but this time it was my own dumb fault: poppy has a relatively small disk and I ran out of space. I fixed that and kicked it off again from the most recent checkpoint, and this time it completed: 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 I converted it to PyTorch for the smoke test: Every effort moves you through each phase, so it's not a complete picture. I'm sure your story was ...which looked solid, so it was time for the test loss eval: 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 Our worst score yet in this experiment! Worse than any of my models so far, apart from the two FineWeb-Edu ones I did without optimised hyperparameters. Now, the first draft of this post went straight to the results from here, but the story wasn't quite over yet... Test set contamination GPT-6 Astra is relentless. Before I publish any of these posts, I run them past an editorial board of LLMs to look for issues. GPT-6 Astra not only checked the text, it also visited the code I'd linked to to check that out too, and spotted something problematic. It's obvious in retrospect, but my code to build the new datasets had a high risk of including the contents of the -- in theory held-back -- test set. The way that the test set was generated was that I downloaded the 10B sample of FineWeb back in December, splitting it into 99% training data and 1% "validation". That validation split was about 100M tokens, and I was only using the first 19M or so for actual validation runs during training, so I (somewhat arbitrarily) designated about 19M other tokens starting at position 50M in there as my test set. Now, my new dataset-generation code was just sampling randomly from the complete 10B sample of FineWeb. So there was nothing stopping it from pulling in data that was in that old validation split! That meant that it was quite likely that my new "curated" and "50:50" datasets contained at least some of the test set that was meant to have been held back from the models during training. On reflection, the problem was potentially even worse. FineWeb-Edu is a subset of FineWeb; my existing FineWeb-Edu dataset came from the 10B sample of the Hugging Face original, and so it also could potentially contain documents that I'd put into the test set. The first thing to do was to establish the size of the problem. I wrote a script to take in a "forbidden" dataset and split; this was assumed to be formatted as one big tensor of GPT-2 tokens, which is what all of my datasets are. It would then split it by end-of-text tokens, and generate a hash and a token count for each resulting "document". Optionally, you could restrict it to only considering a subset -- the n tokens starting at position p -- and it would then generate hashes/lengths for the documents inside that slice, or that overlapped it at the start or the end. I ran that to generate a list of hashes for the entire validation set -- the validation split of gpjt/fineweb-gpt2-tokens -- and then used a second script to check my various training sets (and the validation set itself) to see how much of a contamination problem there was. I got these results: Dataset Split Contamination with validation set gpjt/fineweb-gpt2-tokens validation 102163003 out of 102163003 tokens (100.00%) gpjt/fineweb-gpt2-tokens train 636166 out of 102163003 tokens (0.62%) gpjt/fineweb-edu-gpt2-tokens train 672189 out of 102163003 tokens (0.66%) gpjt/fw-fwedu-5050-gpt2-tokens train 49224580 out of 102163003 tokens (48.18%) gpjt/fw-fwedu-simplewiki-gpt2-tokens train 44233824 out of 102163003 tokens (43.30%) gpjt/openwebtext-gpt2-tokens train 212 out of 102163003 tokens (0.00%) So: The validation set was 100% "contaminated" with itself, which was a useful sanity check. The training set of gpjt/fineweb-gpt2-tokens had what I felt was a small level of contamination. It was interesting that there was any at all -- I think that must mean that there are some repeated documents in the original dataset, and some of them wound up with copies in both my training and validation splits. The gpjt/fineweb-edu-gpt2-tokens dataset also had what felt like a reassuringly low level of contamination. Both gpjt/fw-fwedu-5050-gpt2-tokens and gpjt/fw-fwedu-simplewiki-gpt2-tokens, however, looked problematic. In both cases, the training datasets had more than 40% of the validation/test set in them. gpjt/openwebtext-gpt2-tokens was, as you'd expect, almost completely uncontaminated. It looks like maybe one document happened to have been picked up by both the OpenWebText and the FineWeb crawls and then included in the bit of FineWeb I was using for validation. However, these numbers -- while scary, at least for the 50:50 and the curated datasets -- were not quite the ones to use. They showed how much of the full validation set showed up in the full training set; what I actually cared about was how much of the test set -- those 19M tokens starting at position 50M in the validation split -- was in the actual subset of the training datasets that I actually trained on -- the first ~3.2B of them. I re-ran the script to generate hashes for just the test set, and then re-ran the contamination-checking script, telling it just to look at the appropriate subset of the training tokens, and got this: Dataset (first 3.2B tokens only) Split Contamination with test set gpjt/fineweb-gpt2-tokens train 26557 out of 19632681 tokens (0.14%) gpjt/fineweb-edu-gpt2-tokens train 32079 out of 19632681 tokens (0.16%) gpjt/fw-fwedu-5050-gpt2-tokens train 2986889 out of 19632681 tokens (15.21%) gpjt/fw-fwedu-simplewiki-gpt2-tokens train 2682430 out of 19632681 tokens (13.66%) gpjt/openwebtext-gpt2-tokens train None It was clear that there was a problem -- certainly with gpjt/fw-fwedu-5050-gpt2-tokens and gpjt/fw-fwedu-simplewiki-gpt2-tokens. They'd seen what felt like a significant amount of the test set while training, so their results on the test loss eval were dubious at best. I decided to train those two models afresh, and see what the result was in terms of loss. If the difference was huge, I'd look into the risks of the (much smaller) contamination of gpjt/fineweb-gpt2-tokens and gpjt/fineweb-edu-gpt2-tokens. But if it was pretty small, I'd not worry about that too much. I extended the script that prepared datasets so that the config file could specify a forbidden_dataset. Any documents in the source datasets that matched forbidden ones would be excluded from the output. I then updated the config for gpjt/fw-fwedu-5050-gpt2-tokens and gpjt/fw-fwedu-simplewiki-gpt2-tokens so that the whole validation split of gpjt/fineweb-gpt2-tokens was forbidden, and re-generated them. You can see the updated datasets here and here. Running the contamination-checker script against them showed that they were clear. I then re-did the full training runs for those models; the uncontaminated version of the 50:50 split model is here, and the curated one is here. And the good news: both of them actually did very slightly better at the test loss eval than their equivalents that had been trained on the contaminated data: Model Contaminated Test loss JAX, FineWeb/FineWeb-Edu 50:50 No 3.449257 JAX, FineWeb/FineWeb-Edu 50:50 Yes 3.462454 JAX, curated No 3.534068 JAX, curated Yes 3.542460 There are a number of possibilities that come to mind; perhaps learning from the test set just doesn't happen with tiny 163M models like this, or perhaps while the contaminated models were learning, the benefit they got from that was outweighed by the data that they got instead of the test set data being in some way better for training purposes, at least in terms of the loss eval. But anyway, I felt that if the effect of seeing more than 10% of the test set data during training was so tiny, then the effect of seeing less than 0.2% -- which is what the FineWeb-Edu model in this set of training runs had, as did all of my other FineWeb-only models from previous experiments -- would be even smaller and I'd disregard it. That was excellent news! I didn't need to start all of my experiments from scratch. For the rest of this post, I will include the numbers and results for the contaminated models as well as the uncontaminated ones -- they're interesting for several reasons -- but for future posts I'll skip the contaminated ones. So -- finally! -- let's start digging into the final results. Results Firstly, I think it's worth taking a look at all of the test loss results in context. Here they are in a table, with the new models in bold: Test loss OpenAI weights: medium 3.231442 JAX, overtrained one long epoch 3.324953 JAX, overtrained two normal epochs 3.326482 JAX, with MHA bias, no dropout 3.418784 JAX, no MHA bias, no dropout 3.420089 JAX, FineWeb/FineWeb-Edu 50:50 (uncontaminated) 3.449257 JAX, FineWeb/FineWeb-Edu 50:50 (contaminated) 3.462454 JAX, no MHA bias, with dropout 3.476802 OpenAI weights: small 3.499677 JAX, curated (uncontaminated) 3.534068 1xrtx3090-stacked-interventions 3.538161 JAX, curated (contaminated) 3.542460 8xa100m40-stacked-interventions-1 3.577761 JAX, FineWeb-Edu 3.632900 Cloud FineWeb, 8x A100 40 GiB 3.673623 1xrtx3090-baseline 3.683835 8xa100m40-baseline 3.691526 Cloud FineWeb, 8x H100 80 GiB 3.724507 Cloud FineWeb, 8x A100 80 GiB 3.729900 Cloud FineWeb, 8x B200 160 GiB 3.771478 Local FineWeb train 3.943522 JAX, openwebtext 4.045255 Local FineWeb-Edu extended train 4.134991 Local FineWeb-Edu train 4.166892 I think there's something very clear here: with the new models, the more FineWeb that was in the training mix, the better the model did on this eval. I think I might have been subconsciously expecting that in the predictions I did before running these experiments, but in retrospect it's so incredibly obvious that I feel silly for not mentioning it explicitly! But that tells us something interesting. From the description in the paper, whatever OpenAI did the GPT-2 training run on, it was not like FineWeb. It was probably more similar to OpenWebText -- and yet, that model was the one that performed the worst on this test eval, so if it is more like OpenWebText, there must be some other factor involved. But moving on for now: how about the IFT test -- the one that kicked off all of this work in the first place? I generated a set of IFT responses for all of the new models, and then ran them (plus responses for all of the other models on that table above) past GPT 5.5, and found that one of my new models was getting quite close to the original GPT-2 small weights! So I did four more runs, so that I could get an average. Here are the results -- the "IFT score" is the average across all five runs of the judge, and the "IFT rank" is based on that. The "IFT epochs" was from the original result-generation script. Test loss IFT epochs IFT score IFT rank OpenAI weights: medium 3.231442 2 42.36 1 JAX, overtrained one long epoch 3.324953 3 18.67 7 JAX, overtrained two normal epochs 3.326482 4 18.71 6 JAX, with MHA bias, no dropout 3.418784 4 17.90 8 JAX, no MHA bias, no dropout 3.420089 5 20.50 4 JAX, FineWeb/FineWeb-Edu 50:50 (uncontaminated) 3.449257 4 17.69 9 JAX, FineWeb/FineWeb-Edu 50:50 (contaminated) 3.462454 4 19.30 5 JAX, no MHA bias, with dropout 3.476802 5 13.02 21 OpenAI weights: small 3.499677 2 25.19 2 JAX, curated (uncontaminated) 3.534068 4 16.63 10 1xrtx3090-stacked-interventions 3.538161 4 13.51 19 JAX, curated (contaminated) 3.542460 4 13.58 18 8xa100m40-stacked-interventions-1 3.577761 4 10.19 24 JAX, FineWeb-Edu 3.632900 4 24.56 3 Cloud FineWeb, 8x A100 40 GiB 3.673623 3 16.59 11 1xrtx3090-baseline 3.683835 4 15.15 12 8xa100m40-baseline 3.691526 3 13.64 16 Cloud FineWeb, 8x H100 80 GiB 3.724507 4 13.59 17 Cloud FineWeb, 8x A100 80 GiB 3.729900 3 10.79 23 Cloud FineWeb, 8x B200 160 GiB 3.771478 4 13.70 15 Local FineWeb train 3.943522 5 11.87 22 JAX, openwebtext 4.045255 4 13.28 20 Local FineWeb-Edu extended train 4.134991 5 14.29 14 Local FineWeb-Edu train 4.166892 5 14.69 13 If you want to see the full numbers, they're below. The number that initially surprised me, and made me decide to do multiple LLM-judge runs was the one for the "JAX, FineWeb-Edu" model. In my first run it came in at 24.35 vs the OpenAI small weights' 24.93 -- so close that I wondered if it might even beat them on a re-run. However, in the further four runs its score was consistently lower than the OpenAI model's, and the gap extended a bit in some. So, was FineWeb-Edu the clear winner here? Perhaps. If you look at the contaminated/uncontaminated pairs, something interesting pops out. For the 50:50 mix, the model trained with the contaminated dataset got 19.30, and the one trained on the uncontaminated one got 17.69 -- a difference of 1.61. For the "curated" dataset, the situation was even more interesting: uncontaminated got 16.63, while contaminated got 13.58, a delta of 3.05 points. Remember, the contamination issue is about whether or not the model saw the held-back test set during training. It was an issue for the test loss that is based on that test set, but is entirely orthogonal to the IFT test. From the IFT perspective, both contaminated and uncontaminated models in each case saw training data that was -- in theory, at least -- essentially the same in terms of quality. Indeed, the uncontaminated run saw almost the same data in the same order as the contaminated one, except that some items were omitted, and then extra ones were added to the end. The purpose of this set of experiments was to see how data quality affected the results on the IFT test set. But in the case of the curated model, something that should be unrelated to data quality changed the results by 3.05 points! If something as simple as changing which data of the same quality the model is trained with can affect the IFT score so drastically, it makes it a bit harder to be certain as to whether or not data quality really had the effect we were looking for. On the other hand, the FineWeb-Edu model came in at 24.56, which is 4.06 points better than the 20.50 that the closest other model got -- more than the 3.05 points we see in difference between the two curated dataset models. And it's worth noting that the model with 20.50 is "JAX, no MHA bias, no dropout", which has a subtly different architecture -- no bias on the output projection of the multi-head attention blocks. A better comparison might be "JAX, with MHA bias, no dropout", which got a score of 17.90, for a whacking great difference of 6.66 points. I think that without doing a very large number of training runs on different datasets with different mixes, each one created with a different seed, it would be hard to work out exactly what is in the noise here and what is not. However, that would cost a lot in terms of time. I think that the best thing here is to chalk this up as a fairly decent indication that FineWeb-Edu improves matters for the IFT eval, but far from a certainty. But it's certainly worth noting that whatever the noise is, it has a range of at least 3.05 points -- and the FineWeb-Edu model is just 0.63 points short of GPT-2 small! So there could well be something there. Of course, we don't know whether that model got (by chance) the best possible balance of FineWeb-Edu tokens, and could never win -- or whether it got a bad balance and would actually beat GPT-2 with a better one. So that's certainly worth keeping in mind. As an aside, the result for the curated dataset really surprised me. I had expected that it would be the best one, simply because it almost certainly contained more facts. I took a look at its answers to the questions -- one possibility that came to mind might be that it would get better responses to questions like "What is the chemical symbol for chlorine" or "Who wrote Pride and Prejudice" than the others, but would fail on less knowledge-based tasks. But it was terrible at fact-based questions too: Name the author of 'Pride and Prejudice'. What is the periodic symbol for chlorine? As I understand it, many real-world training runs do include (often oversampled) amounts of highly educational training data like this model's dataset did. But perhaps the models that I'm training are just too small to be able to make use of the data they gained that way -- maybe doing things this way and expecting good results is like asking six-year-old children to memorise stuff before they've learned enough to be able to make use of it 1. It's worth noting that the GPT-2 small model also failed on those factual questions. Well, anyway: I think we have some useful results here, so let's work out what that means for next steps. Conclusion The results we got in these experiments point in two interesting directions. The perfect connection between the amount of FineWeb in the training set and the result on the (FineWeb-based) test loss eval, while perfectly obvious in retrospect, really does highlight how mysterious it is that the OpenAI small weights do so well on that test. The fact that FineWeb-Edu did well on the IFT test tells us that there does seem to be value in using richer training data -- though the less-spectacular results of the 50:50 mix and the curated one weaken that a bit, as does the indicator of what the noise due to data selection from equivalently high-quality datasets might be. The OpenWebText result I think I'll ignore, given that -- while in theory it should be similar to what OpenAI trained on -- there are no guarantees, and it might differ in non-obvious ways for non-obvious reasons. I think that the right direction to take this going forward is to separate these two angles. I should chase a higher IFT score, and then once I have nailed that down, I should see what (if anything) might allow me to get the resulting model to improve its test score. But I will need to make sure that whatever dataset I use, I use various "mixes" of it -- versions created with different random seeds. In my earlier experiments with overtraining, I did find that it didn't seem to improve the IFT results -- but it did improve the test loss. So perhaps identifying the right combination of other factors to boost the IFT score, then overtraining the result, might help? Of course, my overtraining tests were with FineWeb, so the connection might not hold up as well if the starting model (as seems likely) was trained on a different dataset. Also, while working through the results here, I've come to the conclusion that the set of models I'm using is a bit confusing -- there are now different hyperparameter settings, small architectural differences (the MHA bias thing), dropout settings during the pre-training, and now datasets. I think that's OK for now; I should see this part of this series as more ideation than actually running the proper experiments. But at the end, when I have some solid hypotheses with a reasonable amount of backup, I should start from scratch: a baseline model, then staged interventions to build up to what (hopefully) will be a model as good as GPT-2 small. Anyway, I'll wrap this one up here. I think that the next lever to pull is (perhaps surprisingly) going to be weight tying. I had previously kind of disregarded that as a possibility, but while I was working on this post, something popped into my mind. The OpenAI models were originally trained with weight tying. My codebase does actually support doing it -- but because I got the OpenAI weights I'm using from the code in "Build a Large Language Model (from Scratch)", when I'm running the IFT test, the weights are not actually tied! We load up a model that has separate but identical embedding and output head matrices, and then we fine-tune that. So those two matrices can vary independently during fine-tuning -- to put it another way, while GPT-2 small was pre-trained with 124M parameters, the IFT test is being done on a 163M-parameter version. Does that give them some non-obvious advantage? And would adding weight-tying to my own models help, either with or without the output heads being independent at fine-tuning time? Stay tuned :-) Appendix: all IFT judge runs Here are the numbers for all of the IFT judge runs, included for completeness. You can see that the LLM judge ranks models very consistently between runs, but there is variation -- that is, on some runs it's in what I think of as a "better mood" than others, and if that's the case, it will give better scores -- but it will give them almost consistently between models, so all of the models do better. Note that (unlike the table above) this one is sorted by the average IFT score rather than the test loss. Model Run 1 Run 2 Run 3 Run 4 Run 5 Average OpenAI weights: medium 42.24 42.16 42.95 41.83 42.61 42.36 OpenAI weights: small 24.93 24.96 25.39 25.01 25.66 25.19 JAX, FineWeb-Edu 24.35 24.55 24.3 24.68 24.9 24.56 JAX, no MHA bias, no dropout 20.5 19.9 20.76 21.25 20.07 20.50 JAX, FineWeb/FineWeb-Edu 50:50 (contaminated) 19.16 18.86 19.61 19.17 19.7 19.30 JAX, overtrained two normal epochs 18.47 18.29 19.17 18.69 18.91 18.71 JAX, overtrained one long epoch 18.04 18.71 19.62 18.41 18.57 18.67 JAX, with MHA bias, no dropout 17.49 17.35 18.33 17.73 18.62 17.90 JAX, FineWeb/FineWeb-Edu 50:50 (uncontaminated) 17.37 17.73 17.53 18.01 17.83 17.69 JAX, curated (uncontaminated) 16.77 16.03 17.3 16.08 16.96 16.63 Cloud FineWeb, 8x A100 40 GiB 16.44 16.23 17.14 16.62 16.54 16.59 1xrtx3090-baseline 14.85 15.07 15.19 15.14 15.51 15.15 Local FineWeb-Edu train 14.37 14.23 15.08 14.79 15 14.69 Local FineWeb-Edu extended train 14.4 14.07 13.82 14.56 14.61 14.29 Cloud FineWeb, 8x B200 160 GiB 13.37 13.05 13.85 13.67 14.57 13.70 8xa100m40-baseline 13.64 13.36 13.9 13.32 13.97 13.64 Cloud FineWeb, 8x H100 80 GiB 13.45 13.32 13.6 13.51 14.07 13.59 JAX, curated (contaminated) 13.09 13.48 13.95 13.23 14.15 13.58 1xrtx3090-stacked-interventions 13.37 13.11 14.04 13.84 13.17 13.51 JAX, openwebtext 12.88 12.7 13.74 13.53 13.53 13.28 JAX, no MHA bias, with dropout 13.19 12.86 12.98 12.85 13.24 13.02 Local FineWeb train 11.75 11.75 12.21 11.46 12.19 11.87 Cloud FineWeb, 8x A100 80 GiB 10.68 10.2 11.03 10.55 11.49 10.79 8xa100m40-stacked-interventions-1 9.44 9.79 10.84 10.2 10.66 10.19 A small boy asleep on his right side, the right arm stuck out, the right hand hanging limp over the edge of the bed. Through a round grating in the side of a box a voice speaks softly. "The Nile is the longest river in Africa and the second in length of all the rivers of the globe. Although falling short of the length of the Mississippi-Missouri, the Nile is at the head of all rivers as regards the length of its basin, which extends through 35 degrees of latitude …" At breakfast the next morning, "Tommy," some one says, "do you know which is the longest river in Africa?" A shaking of the head. "But don't you remember something that begins: The Nile is the …" "The - Nile - is - the - longest - river - in - Africa - and - the - second - in - length - of - all - the - rivers - of - the - globe …" The words come rushing out. "Although - falling - short - of …" "Well now, which is the longest river in Africa?" The eyes are blank. "I don't know." "But the Nile, Tommy." "The - Nile - is - the - longest - river - in - Africa - and - second …" "Then which river is the longest, Tommy?" Tommy burst into tears. "I don't know," he howls. Brave New World, Aldous Huxley ↩

3 days ago • 1 votes
Pluralistic: Voting is to politics as shopping is to boycotts (01 Oct 2026)

Today's links Voting is to politics as shopping is to boycotts: The big P only matters if the small p is in play. Hey look at this: Delights to delectate. Object permanence: Gilberto Gil v WIPO; Censored Apple wifi hacker talk; Wells Fargo crime-spree started in 1998; Stencils "may not be reproduced"; DVD Jon v Apple DRM; Unpaid diplomatic parking tickets as index of corruption; Tortured Canadian was not a terrorist; Decarbonization at a distance. Upcoming appearances: Brighton, Virtual, South Bend, Hudson, Calgary, Winnipeg, Paris, Vancouver, Victoria, Ottawa, Kilkenny, Montreal. Recent appearances: Where I've been. Latest books: You keep readin' em, I'll keep writin' 'em. Upcoming books: Like I said, I'll keep writin' 'em. Colophon: All the rest. Voting is to politics as shopping is to boycotts (permalink) Here's a funny thing about the right to vote: it wasn't won by voting. From the Magna Carta to the US Constitution to the Emancipation Proclamation to 19th Amendment, voting rights (what you might call "Big P" Politics) were always downstream of protests, riots, petitions, mass movements, strikes and good, old fashioned community organizing (that is, "small p" politics). Which is to say, Big P politics matter, but to make them matter, we need a lot of small p politics. That means that democracy isn't something you do every couple of years with a ballot paper (though that's an important aspect of the process). Democracy is continuous. If you've ever wondered why your vote seems to accomplish so little, I think you can blame the near-abolition of small p politics by Big P politicians of every stripe. Indeed, Obama's genius was summoning up an army of door-knocking, phone-banking small p political activists and then euthanizing that organization after he won the election: https://newrepublic.com/article/140245/obamas-lost-army-inside-fall-grassroots-machine For Obama, the grassroots were useful for one thing: getting out the vote. The last thing he wanted was for millions of activated voters to turn into activists who'd flame him and harangue him and picket him if they didn't like his compromises. Boy, did Obama ever compromise. He let the bank executives who created the Great Financial Crisis off the hook and encouraged them to foreclose on the homes of millions of Americans, the very same public that had bailed them out: https://theweek.com/articles/624777/obamas-biggest-failure He shielded the CIA's torturers from scrutiny and prosecution: https://journals.law.harvard.edu/ilj/2009/04/obama-publishes-torture-memos-immunizes-cia-staff/ He reneged on his promise to shut down Gitmo: https://www.pbs.org/newshour/show/obama-failed-close-guantanamo And his promise to hold the phone companies to account for their complicity in the NSA's mass domestic surveillance: https://www.pbs.org/wgbh/frontline/article/obama-on-mass-government-surveillance-then-and-now/ He stepped up secret drone warfare: https://www.cfr.org/articles/obamas-final-drone-strike-data And unconstitutional domestic surveillance: https://www.eff.org/deeplinks/2017/01/obama-expands-surveillance-powers-his-way-out Whenever I raise this, Obama's apologists come out of the woodwork to tell me that "the president isn't the Green Lantern," and that Obama couldn't act without help from Congress and the Senate, who wouldn't back his plays. I think that Trump's presidency has shown us how much power the president really has even when the legislature won't play ball. But even if you accept the Green Lantern apologetics, the fact remains that Obama could have had a clamoring army of ardent supporters in the streets, defending his agenda against recalcitrants in his own party and wreckers in the GOP. He chose not to have that army. He sent that army home. It's like Obama heard the story about post-election FDR telling civil rights leaders, "I want to do it, now make me do it," and concluded, "I don't want to do it, so I'd better not let anyone make me do it": https://www.quora.com/Did-Franklin-Roosevelt-ever-say-I-agree-with-you-I-want-to-do-it-now-make-me-do-it Of course, Trump is doing everything he can to extinguish both small p politics and Big P Politics. It's not just his wildly illegal voter suppression tactics. He's banning and prosecuting political groups, invoking anti-terror laws (which Obama supported and promised would only be used proportionately and wisely) to chase his grassroots opposition underground: https://www.whitehouse.gov/presidential-actions/2025/09/designating-antifa-as-a-domestic-terrorist-organization/ Liberals are often contemptuous of grassroots movements (cf "basket of deplorables," "Green Lantern" scolding), but the right is terrified of them. The right's political leadership is terrified of its own grassroots, and rightly so, because those people are maniacs, and they're the reason the GOP has been pushed into its most extreme positions. The right's grassroots, meanwhile, are afraid of the left's grassroots. The last thing they want is a militant, organized, mobilized base pushing Dem politicians to take the stands that are wildly and widely popular in America, from Medicare for All to an end to ICE – the Mamdani agenda, in other words. Mamdani is the anti-Obama. He shows what happens when a progressive candidate nurtures and co-governs with their base after the election, using millions of passionate, committed, everyday people to steamroller anyone who gets in the way of his agenda: https://www.nyc.gov/content/100days/pages/ Of course the downside of this is that when Mamdani reneges on his pledges, he is loudly and furiously held to account for it: https://www.thecityreporter.nyc/2026/02/19/mamdani-budget-parks-libraries/ Mamdani understood that he would be corralled into compromises if he won the mayoralty and that when he made those compromises, his base would come after him with the unmistakable fury of betrayed idealists. He also understood that any comfort he enjoyed by sidelining his base while in office would come at a price far higher than being yelled at by his supporters: it would cost him the ability to get anything done. Voting for Mamdani was important. It got him elected. But staying organized – in unions, neighborhood clubs, affinity groups, DSA chapters and mutual aid groups – is what's letting him get stuff done, and stopping him from bailing on his promises as politically infeasible. In other words, voting only matters if it's the final stage of a sustained campaign to build and mobilize popular power. Without that, voting will get you precious little. The right's leadership understands this very well, which is why they've spent years attacking unions, community organizers like Acorn, and activist institutions like Planned Parenthood. We must defend voting rights – Big P Politics – to the bitter end, but we need to defend organizing – small p politics – just as ferociously. The reduction of politics to voting is part of the 50 year neoliberal project whose foremost goal is to make you think of yourself as an atomized individual and not as a member of a polity. Turning "politics" into "voting" is absolutely in line with Margaret Thatcher's dictum that "there is no such thing as society." It's the same move that convinced workers that the answer to bad working conditions is looking your boss in the eye and threatening to change jobs (not forming a union and striking). It's also the same move that transformed "boycotts" into "shopping." Boycotts are a collective enterprise. Before a boycott takes place, small-p political groups hold meetings, organize alternatives and communicate their demands. During a boycott, organizers work to insulate participants from reprisals, like the Montgomery Bus Boycott organizers who reasoned and remonstrated with employers who disciplined workers whose participation made them late for work. And yes, as part of a boycott, you make some consumption choices. You buy X instead of Y. But "shopping" by itself isn't a boycott. You can't "vote with your wallet" (especially not when billionaires get to vote against you with their wallets): https://pluralistic.net/2025/09/13/consumption-choices/#marginal-benefits Shopping isn't politics, and while voting is Politics (Big P), it's also not politics (small p). A boycott, on the other hand, is politics. What's more, "shopping" has the same relationship to "boycotts" that "voting" has to "politics." It's a step you take, after you've laid a lot of groundwork with other people, as part of a mass movement. I understand why shopping and voting are more attractive than boycotts and politics. Meetings suck. Hell is other people: https://locusmag.com/feature/commentary-cory-doctorow-hell-is-other-people/ But changing the system requires systemic work. Hell is other people because other people are great but it's so hard to get them to do things your way. That takes time and understanding and togetherness and arguing and forgiving. Not everyone has time or capacity for that, and at any given time, we don't all have to be doing that work. We can take turns, spelling each other off at times in our lives when we have more or less slack. But lots of us have to be in the fight, or all of us will get screwed. There aren't enough of us doing politics right now. We can tell, because our politicians are so contemptuous of the grassroots that they will sell us out without a moment's hesitation, smugly certain that they will face no consequences for doing so: https://pluralistic.net/2026/09/22/happy-chudmas/#baloney-in-our-slacks Oligarchs have it easy. Where we have to convince people to fight, they can pay or threaten people to bring them into line. But oligarchs' power is wearing thin. The data-center uprising shows how much fury there is out there, looking for a productive outlet: https://www.bloodinthemachine.com/p/with-the-backlash-to-data-centers Data centers are very bad and very visible, so they make for good targets. But data centers are only the physical extrusion of a vast, brutal, extractive system. The most important way to fight data centers is to take everyone you meet protesting one and organize with them to scare the shit out of "your" politicians so they don't dare compromise on anything. Hey look at this (permalink) The Facebook Fake-out https://www.anildash.com/2026/09/29/facebook-fake-out/ Anatomy Unzipped: John of Arderne’s Sweden Scroll (ca. 1425–35) https://publicdomainreview.org/collection/arderne-scroll/ Inside McDonald’s push to have AI price your Big Mac https://www.reuters.com/business/inside-mcdonalds-push-have-ai-price-your-big-mac-2026-09-29/ what is going on with ceiling fans https://mcmansionhell.com/post/829127919552151552/what-is-going-on-with-ceiling-fans From Shitpost to Bullshit https://www.unpopularfront.news/p/from-shitpost-to-bullshit Object permanence (permalink) #25yrsago GWB's press secretary to media: "watch what you do, watch what you say" https://web.archive.org/web/20010926223602/https://www.whitehouse.gov/news/releases/2001/09/20010926-5.html#BillMaher-Comments#BillMaher-Comments #20yrsago Stencils kit “may not be reproduced in any form” https://web.archive.org/web/20061022000842/http://www.fairuseday.com/index.php/2006/10/01/copyright-is-broken/ #20yrsago DVD Jon selling Apple DRM to Apple’s competitors https://web.archive.org/web/20061004191106/https://featured.gigaom.com/2006/10/02/dvd-jon-fairplays-apple/ #20yrsago Unpaid diplomatic parking tickets as index of national corruption https://web.archive.org/web/20130719065306/https://www.theatlantic.com/magazine/archive/2006/10/primary-sources/305203/ #20yrsago Canadian deported to Syria for torture is cleared https://www.theguardian.com/world/2006/oct/02/worlddispatch #20yrsago Gilberto Gil slams WIPO https://fromgeneva.blogspot.com/2006/09/wipo-general-assembly-impressions-from.html #20yrsago Speech given by censored Apple WiFi hacker at ToorCon https://craphound.com/cache_toorcon_2006.txt #10yrsago Company suspected of blame in Office of Personnel Management breach will help run new clearance agency https://www.reuters.com/article/us-usa-security-background-idUSKCN1202M6/ #10yrsago Wells Fargo started demanding fraud of its employees in 1998; Illinois cuts Wells off from state business https://www.citizen.org/wp-content/uploads/wells-fargo-king-of-cross-sell.pdf #10yrsago Google: if you support Amazon’s Echo, you’re cut off from Google Home and Chromecast https://variety.com/2016/digital/news/google-home-amazon-echo-chromecast-1201874125/ #5yrsago How the IMF loan-sharks the global south https://pluralistic.net/2021/10/02/debt-trap/#global-arm-breakers #1yrago Decarbonization at a distance https://pluralistic.net/2025/10/02/there-goes-the-sun/#carbon-shifting Upcoming appearances (permalink) https://www.epl.ca/blogs/post/elbows-up-with-cory-doctorow/ Brighton: Digital Sovereignty and the Post-American Internet (Green Party Conference), Oct 3 https://www.openrightsgroup.org/events/digital-sovereignty-and-the-post-american-internet/ Virtual: How to govern technology in a multipolar digital world (Connecting Current), Oct 6 https://connectingcurrent.tech/how-to-govern-technology-a-multipolar-digital-world/ South Bend: An Evening With Cory Doctorow (Notre Dame), Oct 6 https://franco.nd.edu/events/2026/10/06/an-evening-with-cory-doctorow/ Hudson, OH: Hudson Library, Oct 7 https://engagedpatrons.org/EventsExtended.cfm?SiteID=3850&amp;EventID=596952&amp;PK= Calgary: Wordfest, Oct 8 https://wordfest.com/2026/show/wordfest-presents-cory-doctorow-2026/ Winnipeg: McNally Robinson, Oct 9 https://www.mcnallyrobinson.com/event-18991/An-Evening-with-Cory-Doctorow Paris: Slow Tech Summit, Oct 15 https://slowtechsummit.com/ Vancouver: Read, Resist, Repair, Rejoice (Vancouver Writers Festival), Oct 19 https://writersfest.bc.ca/festival-event-2026/01 Victoria: Munro's Books, Oct 20 https://www.munrobooks.com/events/6113620261020 Vancouver: Life After AI (Vancouver Writers Festival), Oct 22 https://writersfest.bc.ca/festival-event-2026/46 Ottawa: Life After AI (Ottawa Writers Festival), Oct 24 https://writersfestival.org/event/life-after-ai Kilkenny (Kilkenomics), Nov 6-8 https://kilkenomics.com/ Vancouver: Enshittification (Sid Williams Theatre Society), Nov 10 https://www.sidwilliamstheatre.com/events/cory-doctorow-talks-enshittification/ Vancouver: BC Policy Solutions Gala, Nov 12 https://bcpolicy.ca/gala/ Montreal: World Science Fiction Convention, Sep 2-6 https://montreal2027.ca/en Recent appearances (permalink) Terms of Service with Clare Duffy (CNN) https://www.cnn.com/audio/podcasts/terms-of-service-with-clare-duffy/episodes/458ce968-af5d-11f0-b539-13ed2afe25f8 AI, Work, and Power (Software Engineering Daily) AI, Work, and Power https://softwareengineeringdaily.com/podcasts/cory-doctorow-on-ai-work-and-power/ AI, Corporate Power, and the Fight for Worker Control (Plutopia) https://plutopia.io/cory-doctorow-ai-corporate-power-and-the-fight-for-worker-control/ How to Think About AI—Before It’s Too Late (Daniel Solove) https://www.youtube.com/watch?v=_0xR3uEgGcc Could Tech Bosses Destroy Life As We Know It? (Politics JOE) https://www.youtube.com/watch?v=PL4VktU0SgY Latest books (permalink) "The Reverse-Centaur's Guide to AI," a short book about being a better AI critic, Farrar, Straus and Giroux, June 2026 https://us.macmillan.com/books/9780374621568/thereversecentaursguidetolifeafterai/ "Canny Valley": A limited edition collection of the collages I create for Pluralistic, self-published, September 2025 https://pluralistic.net/2025/09/04/illustrious/#chairman-bruce "Enshittification: Why Everything Suddenly Got Worse and What to Do About It," Farrar, Straus, Giroux, October 7 2025 https://us.macmillan.com/books/9780374619329/enshittification/ "Picks and Shovels": a sequel to "Red Team Blues," about the heroic era of the PC, Tor Books (US), Head of Zeus (UK), February 2025 (https://us.macmillan.com/books/9781250865908/picksandshovels). "The Bezzle": a sequel to "Red Team Blues," about prison-tech and other grifts, Tor Books (US), Head of Zeus (UK), February 2024 (thebezzle.org). "The Lost Cause:" a solarpunk novel of hope in the climate emergency, Tor Books (US), Head of Zeus (UK), November 2023 (http://lost-cause.org). "The Internet Con": A nonfiction book about interoperability and Big Tech (Verso) September 2023 (http://seizethemeansofcomputation.org). Signed copies at Book Soup (https://www.booksoup.com/book/9781804291245). "Red Team Blues": "A grabby, compulsive thriller that will leave you knowing more about how the world works than you did before." Tor Books http://redteamblues.com. "Chokepoint Capitalism: How to Beat Big Tech, Tame Big Content, and Get Artists Paid, with Rebecca Giblin", on how to unrig the markets for creative labor, Beacon Press/Scribe 2022 https://chokepointcapitalism.com Upcoming books (permalink) "The Post-American Internet," a geopolitical sequel of sorts to Enshittification, Farrar, Straus and Giroux, 2027 "Unauthorized Bread": a middle-grades graphic novel adapted from my novella about refugees, toasters and DRM, FirstSecond, April 20, 2027 "Enshittification, Why Everything Suddenly Got Worse and What to Do About It" (the graphic novel), Firstsecond, 2027 "The Memex Method," Farrar, Straus, Giroux, 2027 Colophon (permalink) Today's top sources: Currently writing: “Once Is Enemy Action,” a science fiction novel about the origins of modern technofascism. Today's words: 509 (20770 total). "The Post-American Internet," a sequel to "Enshittification," about the better world the rest of us get to have now that Trump has torched America. Fourth draft completed. Submitted to editor. A Little Brother short story about DIY insulin PLANNING This work – excluding any serialized fiction – is licensed under a Creative Commons Attribution 4.0 license. That means you can use it any way you like, including commercially, provided that you attribute it to me, Cory Doctorow, and include a link to pluralistic.net. https://creativecommons.org/licenses/by/4.0/ Quotations and images are not included in this license; they are included either under a limitation or exception to copyright, or on the basis of a separate license. Please exercise caution. How to get Pluralistic: Blog (no ads, tracking, or data-collection): Pluralistic.net Newsletter (no ads, tracking, or data-collection): https://pluralistic.net/plura-list Mastodon (no ads, tracking, or data-collection): https://mamot.fr/@pluralistic Bluesky (no ads, possible tracking and data-collection): https://bsky.app/profile/doctorow.pluralistic.net Medium (no ads, paywalled): https://doctorow.medium.com/ Tumblr (mass-scale, unrestricted, third-party surveillance and advertising): https://mostlysignssomeportents.tumblr.com/tagged/pluralistic "When life gives you SARS, you make sarsaparilla" -Joey "Accordion Guy" DeVilla READ CAREFULLY: By reading this, you agree, on behalf of your employer, to release me from all obligations and waivers arising from any and all NON-NEGOTIATED agreements, licenses, terms-of-service, shrinkwrap, clickwrap, browsewrap, confidentiality, non-disclosure, non-compete and acceptable use policies ("BOGUS AGREEMENTS") that I have entered into with your employer, its partners, licensors, agents and assigns, in perpetuity, without prejudice to my ongoing rights and privileges. You further represent that you have the authority to release me from any BOGUS AGREEMENTS on behalf of your employer. ISSN: 3066-764X

3 days ago • 1 votes
A big-tent or small-tent AI safety movement?

The unstated disagreement that underpins safety debates

3 days ago • 1 votes
📚 BoredReading

You seem to be enjoying this.

Join free to unlock everything.

Create free account

Already have an account? Sign in