lab

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:

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:

  1. load_file is eager. Even with backend="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 with safe_open(...).get_tensor() costs one tensor instead of the file.
  2. The random model is a whole extra copy. Built with nnx.eval_shape it costs nothing above the interpreter baseline instead of 1.06× the model.
  3. 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

What I would tell anyone writing their own state loader

  1. 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_dict validates extra keys and ignores absent ones, which is exactly backwards for this purpose.
  2. 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.
  3. 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.
  4. 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.
  5. 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.
  6. 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.
  7. 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.