import orbax.checkpoint as ocp from flax import nnx # Manager Setup mngr = ocp.CheckpointManager( directory, options=ocp.CheckpointManagerOptions(max_to_keep=3, save_interval_steps=100), ) # Saving a Single Pytree state = nnx.state(model) mngr.save(step, args=ocp.args.StandardSave(state)) mngr.wait_until_finished() # Restoring a Single Pytree abstract_model = nnx.eval_shape(lambda: MyModel(rngs=nnx.Rngs(0))) graphdef, abstract_state = nnx.split(abstract_model) restored_state = mngr.restore( mngr.latest_step(), args=ocp.args.StandardRestore(abstract_state) ) restored_model = nnx.merge(graphdef, restored_state) # Composite Save/Restore (model + optimizer) mngr.save(step, args=ocp.args.Composite( params=ocp.args.StandardSave(params_state), optimizer=ocp.args.StandardSave(optimizer_state), )) restored = mngr.restore(step, args=ocp.args.Composite( params=ocp.args.StandardRestore(abs_params_state), optimizer=ocp.args.StandardRestore(abs_optimizer_state), )) # Updating Live Objects nnx.update(model, restored["params"]) nnx.update(optimizer, restored["optimizer"]) # Cleanup mngr.close() __ __