import orbax.checkpoint as ocp from flax import nnx # Assume model is an initialized nnx.Module mngr = ocp.CheckpointManager("/tmp/my_checkpoints") # Extract the state state_to_save = nnx.state(model) # Save at step 100 mngr.save(100, args=ocp.args.StandardSave(state_to_save)) # IMPORTANT: wait for the (possibly asynchronous) save to actually finish mngr.wait_until_finished() mngr.close() __ __