nnx_save — a year later: four blocking bugs, one silent one, and a 4× memory tax
nnx_save is a small library I wrote in March 2025 for
exactly one job: save and load Flax NNX models as a
single .safetensors file. It exists because I was porting
models from PyTorch to JAX for TPU inference and did not want to stand
up orbax for something whose whole shape I already had in
my head — flatten the state, write it, read it back, replace the state.
The README apologises for the library; the honest reason it exists is
that I wanted a two-function checkpoint path and did not want to write
one again for every port.
What I never did was test it properly. This report is what happened when I did: a 73-test suite across three architecture families, an 8-device sharding simulation, real HuggingFace checkpoints ported by hand, and a memory benchmark run on a CPU box and on the RTX 3050 (over Tailscale) that the port was ultimately for. Fixing the bugs exposed a second problem — the loader wanted about four times the checkpoint in host RAM — which mattered more than the bugs did for the model I actually wanted to port.
Bottom line: the committed version could not save a single
model on the current stack. jax 0.11.1,
flax 0.12.9 and safetensors 0.8.0 have moved
since March 2025, and the first tensor save_model touched
raised
AttributeError: 'jaxlib._jax.ArrayImpl' object has no attribute 'ctypes'.
One line below that crash sat something worse: a model built with
nnx.List — which is every transformer stack — loaded with
random weights and reported success. That is the
failure mode this whole exercise was really about, because it is the one
that would have shipped. After four blocking fixes and five closed
silent-failure modes, the same checkpoint round-trips bit-for-bit in
fp32 and bf16, ported GPT-2, Llama-style and ViT-style models match
their PyTorch logits before and after saving, and the peak host RAM for
a load fell from 3.0× the checkpoint to 1.1× — measured
on the real tensor shapes of a 2.27 GB Stable Audio 3 checkpoint.
What the library does
Two functions, one file format:
from nnx_save import save_model, load_model
model = nnx.Linear(768, 64, rngs=nnx.Rngs(0))
save_model(model, "model.safetensors") # flat dict of tensors
loaded, state = load_model(model, "model.safetensors") # same structure required
Saving is nnx.state(model) →
nnx.to_pure_dict → flatten to "a/b/c" keys →
safetensors.numpy.save_file. Loading is the reverse:
re-nest the file, keep the values whose paths exist in the model you
pass in, nnx.replace_by_pure_dict,
nnx.merge.
The deliberate limitation is that only the state is stored, never the graph definition. The caller rebuilds the model and the loader fills it in. That is what makes the file portable — and it is also why the loader's job is entirely about matching file keys to model paths, which is where every bug below lives.
How it was tested
| component | version |
|---|---|
| Python | 3.13.15 |
| jax / jaxlib | 0.11.1 |
| flax | 0.12.9 |
| safetensors | 0.8.0 |
| torch (porting reference) | 2.14.0+cpu on CPU; 2.13.0+cu130 on the 3050 |
| machines | 8-core / 31 GB CPU box; victoria, RTX 3050 6 GB over
Tailscale |
Nine suites, 73 tests: 11 round-trip tests (models, containers,
dtypes), 12 failure-mode tests, 13 streaming tests (format validity,
equivalence with the whole-file path, the abstract-builder entry point,
header layout), 6 TPU-style tests, 12 pod-layout tests (per-process
shard directories and single-file shard ranges), and 19 port tests
across GPT-2, a Llama-style decoder and a ViT-style encoder.
tests/conftest.py sets
XLA_FLAGS=--xla_force_host_platform_device_count=8, so
parameters really are sharded across 8 simulated devices — no TPU was
available, and that is stated wherever it matters. The baseline run was
taken before touching the library, against the
committed code.
It could not save anything
$ pytest tests/ -q
27 failed, 1 passed, 2 errors
Every failure is one line inside safetensors:
safetensors/numpy.py:22: AttributeError: 'jaxlib._jax.ArrayImpl' object has no attribute 'ctypes'
safetensors.numpy.save_file builds its tensor table from
tensor.ctypes.data, tensor.dtype.byteorder and
tensor.nbytes: it wants real np.ndarray
objects. nnx.to_pure_dict hands back jax
arrays. So the very first thing save_model does
with a parameter fails:
from flax import nnx
from checkpointer import save_model
model = nnx.Linear(4, 2, rngs=nnx.Rngs(0))
save_model(model, "model.safetensors")
# AttributeError: 'jaxlib._jax.ArrayImpl' object has no attribute 'ctypes'
Two fixes, both one-liners: np.asarray each leaf on the
way out, and nnx.to_pure_dict(state) instead of the
now-deprecated state.to_pure_dict() (flax 0.12 warns that
nnx.State is being replaced by plain dicts, so that
spelling has an expiry date).
The bug that would have cost a week
Patch only the numpy conversion and the library appears to work. It saves, it loads, a plain MLP comes back bit-identical. Then you point it at a transformer:
class Deep(nnx.Module):
def __init__(self, *, rngs):
self.layers = nnx.List([nnx.Linear(4, 4, rngs=rngs) for _ in range(3)])
deep = Deep(rngs=nnx.Rngs(0))
save_model(deep, "deep.safetensors")
loaded, _ = load_model(Deep(rngs=nnx.Rngs(99)), "deep.safetensors")
# per-layer weights restored: [False, False, False]
No exception. No warning. The saved file looks right —
layers/0/kernel, layers/1/kernel, … — and
every layer comes back at its random initialisation.
The cause is a type mismatch hiding behind string formatting.
nnx.List state keys are integers;
safetensors keys are always strings. Saving joins the
path into "layers/0/kernel"; loading splits it back into
the nested dict {"layers": {"0": …}}; the model's own state
has {"layers": {0: …}}. The loader flattened both sides and
intersected the key sets, so ('layers', 0, 'kernel') never
matched ('layers', '0', 'kernel'), and the intersection was
empty for every list-backed parameter.
Why was it silent, when passing a garbage dict to
nnx.replace_by_pure_dict normally raises? Because that
function is asymmetric. It validates keys that are present in the pure
dict but missing from the state
(ValueError: key in pure_dict not available in state), and
says nothing about state entries with no corresponding key in the pure
dict. "Filter the checkpoint down to what the model has" therefore turns
any naming mismatch into a model that is half-random and completely
confident.
For an inference port this is the worst possible outcome: a 12-layer stack loads in a second, the script prints no errors, and the wrongness only shows up later as "the port produces garbage". The fix is to stop intersecting flattened keys and walk the model's structure instead, pulling values out of the checkpoint by matching each expected key against both its native and string form. That also makes the loader tolerant of the files the old implementation wrote.
RNG state could not be saved at all
The third blocker shows up the moment a model contains
nnx.Dropout, nnx.Rngs or anything else that
keeps a generator:
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.
nnx.state includes the Rngs keys and
counters, and a typed key (key<fry>) has no numpy
representation — deliberately. The fix is the one jax suggests: store
jax.random.key_data(key) (plain uint32 bits)
and rebuild with jax.random.wrap_key_data(data, impl=...)
on load, so a resumed training run continues the same random stream
instead of restarting it.
The hazards that were still silent
Fixing the crashes exposed the more interesting question: what else
can a checkpoint loader get wrong without saying so? Five cases, each
now reported rather than swallowed — with strict=True
turning the warnings into errors:
| what goes wrong | before | now |
|---|---|---|
| parameter missing from the file | random weights kept, silence | warning naming the paths; strict=True raises |
| saved array has the wrong shape | wrong-shaped array installed | skipped, warning with both shapes; never installed |
| checkpoint dtype ≠ model dtype | f32 checkpoint silently turns a bf16 model into f32 | cast to the dtype the model declares (PyTorch
load_state_dict behaviour), with a warning |
| values come back as host numpy | parameters off-device, sharding gone | jax.device_put onto the target variable's sharding — a
sharded bf16 model stays sharded |
| checkpoint has keys the model never uses | ignored | warning (this is how a wrong port mapping surfaces) |
Two smaller ones came out of the same pass. Python scalar variables
(nnx.Variable(0) step counters) used to come back as 0-dim
arrays — a silent type change in the middle of a training loop — and now
come back as the type the model declared. String values and
/ in an attribute name cannot be represented in flat
safetensors keys at all; both now fail at save time
with a message naming the offending parameter, instead of producing a
file that misloads later.
The sharding line is the one that matters for the original use case.
Wrapping every loaded value in jnp.asarray is enough to
make the tests pass, but it also quietly collapses a model that was
sharded across 8 devices onto one — which on a real TPU means a model
that no longer fits. The test for it is blunt: shard an MLP across 8
simulated devices, save it, load it into a freshly sharded model, and
assert every parameter is still spread over 8 devices.
How the current version is built
Two modules, ~850 lines, no new dependencies:
| file | what lives there |
|---|---|
nnx_save/checkpointer.py |
one machine, one file: save_model,
load_model, the flatten/match logic, the streaming reader
and writer, the safety report |
nnx_save/sharded.py |
many processes: save_sharded,
load_sharded, the manifest and per-process shard files, the
single-file shard-range reader |
Saving. The model state goes through
nnx.to_pure_dict and is flattened into "a/b/c"
keys (a / inside an attribute name is refused up front,
because it would be ambiguous in the file). Each leaf becomes a
safetensors tensor: jax arrays are viewed as numpy, typed PRNG keys are
stored as their uint32 key data, python scalars become 0-d
arrays. The writer then emits the format by hand — 8-byte header length,
JSON header padded to an 8-byte boundary, then the tensors in sorted key
order — one tensor at a time, so only one host buffer is alive.
(safetensors.save() builds the whole file in memory first,
and save_file requires every source buffer to stay alive
for the duration of the call.)
Loading. The walk is driven by the
model, never by the file. The model is split into graph
definition and abstract state, and every leaf it expects is looked up in
the checkpoint by its own path — accepting both the native and the
string form of each key, which is what lets nnx.List
indices (integers) line up with safetensors keys (always strings).
Values arrive one tensor at a time from
safe_open(..., backend="pread"), are cast to the dtype the
model declares on the host, and are placed with
jax.device_put(..., sharding) so a sharded model stays
sharded. Anything that cannot be applied exactly — a parameter the file
does not have, a shape that does not match, a dtype that had to be cast,
a key the model never uses — is collected and reported: warnings by
default, errors under strict=True.
Into a builder. load_model also accepts
a zero-argument callable, which it builds with
nnx.eval_shape. The parameters exist as
jax.ShapeDtypeStruct shapes until the checkpoint fills them
in, so the "randomly initialised model that gets overwritten" never
exists at all. The one thing lost on that path: an abstract build cannot
tell a python int from a 0-dim array, so scalar counters
come back as 0-d arrays.
Many processes. save_sharded writes a
directory — a manifest plus one shard_NNNNN.bin (and a
sidecar) per process. Saving takes this process's slice of
every global array with
multihost_utils.global_array_to_host_local_array and merges
the sidecars after a barrier; loading reads only the local file and
assembles globals with host_local_array_to_global_array,
one tensor at a time. load_sharded also accepts a plain
.safetensors file: each local device then reads only its
own index range through safetensors' slice API and the array is
assembled with jax.make_array_from_single_device_arrays.
Both directions refuse a sharding this process cannot address, because a
single-file device_put (or np.asarray) would
otherwise write or read just the local slice while reporting
success.
The rules that fell out of the exercise, and that the code now enforces: drive from the model's structure, one tensor at a time, adopt the model's dtype and sharding, and never let a mismatch pass in silence.
The port that motivated the library
The point of all this was never a linear layer. To check the actual
workflow, I downloaded hf-internal-testing/tiny-random-gpt2
(64 tensors, tied embeddings, 5 layers), computed reference logits with
PyTorch on the RTX 3050, and then ported the checkpoint into an NNX
GPT-2 in a JAX-only environment — two processes, two
virtualenvs, so nothing about the port could lean on torch being
importable.
config: GPT2Config(vocab_size=1000, n_positions=512, n_embd=32, n_layer=5, n_head=4, layer_norm_epsilon=1e-05)
checkpoint tensors: 64, tied=True
ported 64 tensors, 0 unknown keys: []
[port] max|jax-torch| = 3.787e-06 (mean 6.161e-07)
[save] 0.5 MB in 0.01s
[load] 0.33s
[after] max|after-before| = 0.000e+00
[bf16] dtypes preserved, max relative deviation from fp32 torch = 3.805e-03
ALL CHECKS PASSED
3.8e-06 is ordinary fp32 reassociation noise between two
independent implementations, and it is the number that says the port is
correct. The 0.000e+00 after the round-trip says
the checkpoint path adds nothing of its own: the loaded model produces
bit-identical logits to the model that was saved, and still matches
torch.
Two porting details are worth writing down, because both cost me time:
- HF GPT-2's
Conv1Dweights are already(in, out), the same layout as an NNXnnx.Linearkernel, so the projection weights copy across without a transpose. Most other HuggingFace architectures usenn.Linear((out, in)) and do need one; getting this backwards is exactly the kind of error the wrong-shape check now catches. - A tied parameter lives at exactly one path, and which one is
an NNX traversal detail. With
self.lm_head = self.wte,nnx.statereports the shared embedding aslm_head/embedding, notwte/embedding. A port that hard-codes the name it expects getsValueError: key in pure_dict not available in state: ('wte', 'embedding'). Resolve the target path against the model's actual state rather than the name in the PyTorch checkpoint.
One GPT-2 was not enough to believe any of this, so the same exercise now covers two more families, each with a torch reference and the same checks:
| arch | what it adds | max abs difference, jax vs torch (fp32) |
|---|---|---|
| GPT-2 (real HF checkpoint) | tied embeddings, Conv1D layout |
3.79e-06 |
| Llama-style (real HF checkpoint) | RMSNorm without bias, HF-convention RoPE, GQA (4 query / 2 KV heads), SwiGLU, no biases, untied head | 8.94e-08 |
| ViT-style | patch-embedding conv, class token, position embeddings, pre-LN blocks | 3.35e-08 |
Each port has a per-tensor coverage test — every checkpoint tensor must land on a model parameter and compare equal, which is what catches a missing or transposed tensor that a bare round-trip check cannot see — plus bit-identical fp32/bf16 round-trips and a load through a builder.
The 4× memory tax
The benchmark I ran after the fixes said a 503 MB checkpoint needed 2.16 GB of host RAM to load. Instrumenting every phase, each strategy in its own process, showed where it went on a 125.8 M-parameter fp32 model:
| phase | RSS |
|---|---|
| interpreter + jax, before any model | 188 MB |
| + a randomly initialised model | 721 MB (+533) |
+ safetensors.numpy.load_file of the whole file |
1343 MB (+622) |
| + the jax host copies and the merge | 1680 MB peak (3.34× the model) |
Three of those terms are avoidable, and the library now avoids them:
load_fileis eager. Even withbackend="mmap"the numpy binding copies the whole file into anonymous host RAM — unlike the torch binding, which aliases the mapping. Reading one tensor at a time withsafe_open(...).get_tensor()costs one tensor instead of the file.- The random model is a whole extra copy. Built with
nnx.eval_shapeit costs nothing above the interpreter baseline instead of 1.06× the model. - Casting on the host, before the transfer, keeps the device side to one temporary rather than two.
Measured end to end on that model: 3.34× → 1.58× peak host RAM, and the load itself about twice as fast.
At the size that actually matters
"Fine for a toy" is not "fits". The real target is the Stable Audio 3
port, so the benchmark now takes the exact tensor
shapes from a real checkpoint header —
stabilityai/stable-audio-3-small-music: 685 tensors, 567.6
M parameters, 2.27 GB fp32, largest tensor 8.4 M parameters, median 1024
— and builds synthetic parameters of those shapes. No weights
downloaded.
| load path | 8-core CPU box | victoria (7.5 GB RAM) |
|---|---|---|
| whole file into memory, random model (the original) | 6.87 GB (3.03×), 5.7 s | does not fit |
| streaming, random model | 4.74 GB (2.09×), 4.3 s | — |
| streaming + builder (the new default) | 2.56 GB (1.13×), 4.0 s | 2.49 GB (1.10×), 1.9 s |
That is the difference between "cannot load the full SA3-small
checkpoint on the small host that runs its inference" and "loads it in
2.5 GB, in about two seconds". bf16 halves the weights again. The same
arithmetic on the 9.22 GB medium checkpoint (which I did
not download) puts the old path near 30 GB of host RAM and the new one
near 11 GB.
One caveat the numbers make visible: 685 small tensors cost more per-tensor overhead than 135 large ones of the same total volume, which load in 1.7 s. Per-tensor streaming trades a little time for a lot of memory — the right trade when the alternative is not loading at all.
TPU-resident checkpoints
A single file has a pod-shaped problem:
np.asarray(global_array) gathers every host's shards into
one process, and
jax.device_put(full_array, global_sharding) requires
every host to hold the whole array (jax asserts the inputs
match across hosts, then defers the shard). So the library writes the
layout pod checkpointers use when you ask for it:
ckpt_dir/
manifest.json # dtype, global shape, PartitionSpec, per-process offsets
shard_00000.bin # process 0's local shards, concatenated
shard_00000.json # process 0's sidecar, merged into the manifest
shard_00001.bin
Each process writes and reads only its own file, through
global_array_to_host_local_array and
host_local_array_to_global_array, with the sidecars merged
after a sync_global_devices barrier. If you would rather
keep one file, load_sharded accepts a
.safetensors too: each local device reads only its own
index range through safetensors' slice API and the global array is
assembled with jax.make_array_from_single_device_arrays —
the same idea as orbax v1's SafetensorsLayout, without
giving up the single-file format.
Verified on 8 simulated CPU devices: values equal the single-file path, every parameter comes back sharded, a replicated bias stays replicated, a bf16 tied-embedding model round-trips, and damage is still reported. Not verified: the multi-process handshake. There is no TPU here, and that is the honest boundary of this work.
How the established libraries avoid the tax
I read the installed orbax (0.12.4), jax, safetensors and torch sources rather than guessing; the mechanisms line up with what is now in nnx_save:
| mechanism | who | effect |
|---|---|---|
one tensor at a time, never load_file |
orbax, HuggingFace, this library | ~1× the model |
abstract/meta init (nnx.eval_shape,
torch.device("meta")) |
orbax v1 restore targets; transformers, unconditionally
since v4.51 — low_cpu_mem_usage no longer exists |
~1× |
| read only this process's shards | orbax reads sharding._addressable_device_assignment,
assembles with make_array_from_single_device_arrays |
1/n_hosts |
| per-process shard directories | orbax ocdbt.process_N; PyTorch/XLA DCP
__<rank>_<n>.distcp |
standard pod layout |
| ranged reads in one file, driven by the target sharding | orbax v1 SafetensorsLayout — "no cross-process
communication and no XLA compilation" |
per-host bytes |
| an in-flight byte budget | orbax restore_concurrent_bytes /
MemoryOptions.read_concurrent_bytes (2 GiB, 128 MiB
chunks) |
O(budget) |
| bind instead of copy | load_state_dict(assign=True), accelerate's
set_module_tensor_to_device, transformers'
setattr |
one destination allocation |
Two corrections worth recording, because both are easy to get wrong
in a write-up: orbax has no TpuShardedArrayHandler
or JaxArrayHandler — it is one
ArrayHandler plus SingleReplicaArrayHandler;
and there is no TPU-direct-to-GCS write path — device
memory goes to host RAM and then through TensorStore, with a RAM disk as
the recommended staging tier on TPU VMs (the genuine device-side path is
Pathways-on-Cloud Persistence).
Still not handled
- State only. No graph definition is stored, so the caller must rebuild the same structure. A structural mistake is now visible (missing/extra-key warnings) but not fatal by default.
- Sharding is adopted, not stored. Load into a model built with the sharding you want. Load into an unsharded model and you get replicated parameters — deliberately, but it is a decision the caller has to make.
- Pod coordination is unverified. The sharded paths are written for a multi-process pod and refuse anything that would silently touch only the local slice, but no pod was available to run the handshake — that is the first thing to try on real hardware.
- No in-flight byte budget. Streaming bounds the peak
at one tensor; a single multi-gigabyte tensor would still be a spike.
orbax's
restore_concurrent_bytesequivalent is the fix, and it matters above ~10 GB. - No async overlap. The device→host copy is blocking,
so a save stalls the caller; orbax's
AsyncCheckpointerkeeps the copy blocking but moves the write off the critical path. - Saving is not atomic. A crash
mid-
save_fileleaves a partial checkpoint; write to a temp path andos.replaceif that matters. - CUDA JAX was not exercised. The 3050 ran the torch
half of the port; the jax half ran on CPU. Getting CUDA jax onto that
host needed ~3.5 GB of
nvidia-*wheels over a ~1 MB/s link (the plugin-only shortcut installed but never registered a backend), so it was left for a session with a faster pipe.
What I would tell anyone writing their own state loader
- Silent is worse than broken. A load that reports
success while keeping random weights is the only outcome that can
survive to production. Check for missing keys yourself:
replace_by_pure_dictvalidates extra keys and ignores absent ones, which is exactly backwards for this purpose. - Poison the model before loading it in your tests. Overwrite every float parameter with a constant, then load — otherwise a no-op loader passes your test suite.
- Round-trip every container and dtype you actually
use.
nnx.List, tied parameters, rng state, batchnorm buffers, bf16, scalar counters. The container was the one that broke here, and it is the one people forget. - Check the port against the source framework, not against yourself. A round-trip test proves only that two copies of the same model agree; it cannot tell you whether the mapping from PyTorch was right. The comparison against torch logits is what makes the port claim mean something.
- Stream the tensors, and never build a model you are about to overwrite. Whole-file loading and a random initialisation are each about one extra copy of the model in host RAM; on a 2.27 GB checkpoint that is the difference between fitting and not fitting.
- A benchmark that keeps both models alive lies to you. The first version of my own scale benchmark held the saved-from model and the fresh one, and reported 3.9× where the honest single-model figure was 2.1×. Measure the phase, not the vibe.
- Assume API drift. This library was fine for the
safetensors of early 2025 and useless for the safetensors of 2026
because one internal detail changed from
.tobytes()to.ctypes.data. Version-pin your checkpointing path, or test it on a schedule.
Reproduce
# whole suite on CPU (73 tests, ~100 s)
.venv/bin/python -m pytest tests/ -q
# the real checkpoint, two environments: torch reference, then the jax port
# (first command needs torch + transformers, second needs jax + flax only)
python tests/port_hf_reference.py --arch llama --out /tmp/hf_llama
python tests/port_hf_check.py --dir /tmp/hf_llama
# host RAM per load strategy, on a synthetic model or on real checkpoint shapes
.venv/bin/python tests/bench_memory_attribution.py --params-millions 568
.venv/bin/python tests/bench_memory_attribution.py --shapes-from sa3_shapes.json
The library, the suite and the full evidence table are in the
repository: nnx_save/checkpointer.py (483 lines),
nnx_save/sharded.py (368 lines) — up from 90 lines in one
file — and tests/REPORT.md.