Writing an LLM from scratch, part 34a -- building a JAX training loop for an LLM training run
For over a year, I've been using Sebastian Raschka's book "Build a Large Language Model (from Scratch)" -- and the multitude of side-projects that have branched out from reading it -- as something like a curriculum for learning about modern AI. The one final task I had set myself was 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 following the book.
As an output, I wanted something as good as my best PyTorch model based on Raschka's code -- a base model, trained on 3.2B tokens, that my (admittedly limited) evals ranked as being close to the original GPT-2 small's quality.
I wanted to use a different framework, just to make sure I wasn't parroting code that I'd somehow memorised, so I asked people on Twitter which one I should use, and the winner was JAX.
I took a slightly different route to Raschka's book; he takes an inside-out perspective, explaining things like attention, gradually building up a complete GPT-2-style model, and then building a training loop on top of it. I wanted to go outside-in: I'd put together a training harness to train the simplest-possible model with an API similar to a real LLM, get that working to my satisfaction, and then add features to that simple model, one by one, until it had the full architecture in place. The plan (which actually worked out nicely!) was that I'd be able to show how each change improved things.
That's all done now, and I'm posting about it in two parts; in this one, I'll explain how I built the training harness, and in the next, I'll show the actual building and training of the LLM.
So let's get started!
Thoughts on Role Confusion
The other day, I came across "Prompt Injection as Role Confusion"
(via Simon Willison). It's a really
interesting blog-style version of a paper by Charles Ye, Jasmine Cui and Dylan Hadfield-Menell,
where they find that LLMs seem to almost ignore 'role' tags like <system>, <user> or <think>, and
instead use the tone of text to infer roles. This seems to explain a lot of jailbreaks.
Flax debugging: making a hash of things
I was debugging an issue with a JAX/Flax NNX training loop the other day, and found a neat little trick to help debug it. Specifically, I wanted to see if the issue was with my model, my loss function, my optimiser settings, or the "plumbing" of the training loop itself -- were gradients actually coming through and being applied to the parameters?
I could print out the loss and the gradients, but printing out the parameters to see if they were changing was unhelpful -- any given update might only change a small number of parameters, or might change them such a small amount that I'd not notice -- especially given that the model had 77 million of them!
Let's take a look.
10Gb/s Ethernet: switching to a Broadcom SFP+ module
Back in April, I upgraded my home LAN to 10Gb/s. The in-wall cabling is CAT-6 or similar, so I had to use 10GBASE-T. Now, the router I'm using, and the switch in my study, provide 10Gb/s through SFP+ cages; that meant that they needed 10GBASE-T SFP+ modules in order to connect.
That kind of module is known to run hot -- sometimes too hot to actually work. The
modules in reggie, the router, appeared to be running OK (see the linked post above
for charts), but the one in nigel, the study switch, was a worrying 93C. I tried
sticking some mini-heatsinks on it, which
seemed to help a bit. But the weather got warmer, and eventually the module overheated.
I lost access to the Internet from the study, and checking the metrics showed me this:

You can see that it's "flapping": the temperature gets up to a level where the module shuts itself down for its own protection -- about 95C, I think -- and then when it has recovered, it switches on again, the temperature rises, and the process repeats.
I was able to work around the problem by switching on the air conditioning in the study. But normally I only have it on when I'm in there, and keeping aircon on 24/7 just to keep the network working felt like the wrong solution.
It was time to switch to a more power-efficient SFP+ module.
JAX: commitment issues
Imagine you have JAX code like this, and run it on a machine with CUDA set up:
key = jax.random.key(42)
cpu0 = jax.devices("cpu")[0]
with jax.default_device(cpu0):
array = jax.random.randint(
key,
(530640, 6, 1024),
0, 50_000,
dtype=jax.numpy.uint16
)
array.block_until_ready()
item = array[0]
item.block_until_ready()
We're creating a big array, blocking until it's ready (JAX is asynchronous, so this makes sure that it's actually finished creating it), then getting the first item, and as a belt-and-braces thing making sure that that is ready too. How long do you think those last two lines -- a simple retrieval of a 6 x 1024 array from a larger one -- will take? Some tiny fraction of a second would seem reasonable.
But running it on my machine just now, the answer is a bit of a surprise: just over 5 seconds. And if you try to
get array[1] immediately afterwards, it still takes about 1.2s. Further lookups into
array consistently take more than a second -- so while the larger
initial number might be something to do with setup -- maybe internal stuff being JITted -- that's clearly not the whole story.
Something is making these seemingly-simple array lookups take much longer than you'd
expect them to.
Let's dig into that.
JAX backends and devices
There's nothing like writing your own code with a framework to clarify how things
fit together! Continuing with my port of my PyTorch LLM code to
JAX, I wanted to load up a large dataset:
the 10,248,871,837 16-bit unsigned integers in the train split of
gpjt/fineweb-gpt2-tokens.
That's just over 19GiB of data.
from safetensors.flax import load_file
...
full_dataset = load_file(dataset_dir / f"train.safetensors")["tokens"]
When I ran that, I got a CUDA out-of-memory error:
jax.errors.JaxRuntimeError: RESOURCE_EXHAUSTED: Out of memory while trying to allocate 19.09GiB.
That makes sense! The allocation it was trying to do is exactly the size of the data I was trying to load. I have an RTX 3090 with 24 GiB, but some is already used up by the OS, various apps, and a model that the code creates earlier on.
But in PyTorch land, I was used to things being loaded into RAM by default, and only moved over to the GPU when I asked it to do that. JAX was clearly loading to the GPU by default. How could I stop it from doing that for this case? The load into the GPU was happening inside Safetensors, in code I couldn't directly control.
Understanding how to do it helped me understand a little bit more about JAX.
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.