# Build abstract model and optimizer abs_model = nnx.eval_shape(lambda: CNN(rngs=nnx.Rngs(0))) abs_optimizer = nnx.eval_shape( lambda: nnx.Optimizer(abs_model, optax.adamw(0.001), wrt=nnx.Param) ) graphdef, abs_params_state = nnx.split(abs_model, nnx.Param) abs_optimizer_state = nnx.state(abs_optimizer) # Define what we're restoring into restore_targets = { "params": ocp.args.StandardRestore(abs_params_state), "optimizer": ocp.args.StandardRestore(abs_optimizer_state), } # Restore step = mngr.latest_step() restored = mngr.restore(step, args=ocp.args.Composite(**restore_targets)) # Build real, live instances and update them with restored data model_instance = CNN(rngs=nnx.Rngs(1)) optimizer_instance = nnx.Optimizer(model_instance, optax.adamw(0.001), wrt=nnx.Param) nnx.update(model_instance, restored["params"]) nnx.update(optimizer_instance, restored["optimizer"]) mngr.close() __ __