import optax from flax import nnx tx = optax.adamw(learning_rate=0.001) optimizer = nnx.Optimizer(model, tx=tx, wrt=nnx.Param) # Extract each piece of state separately params_state = nnx.split(optimizer, nnx.Param)[1] optimizer_state = nnx.state(optimizer) save_items = { "params": ocp.args.StandardSave(params_state), "optimizer": ocp.args.StandardSave(optimizer_state), } mngr.save( optimizer.step.value, args=ocp.args.Composite(**save_items), ) mngr.wait_until_finished() __ __