def create_abstract_sharded_target(): abstract_model = nnx.eval_shape(lambda: CNN(rngs=nnx.Rngs(0))) _, abstract_state = nnx.split(abstract_model) sharding_specs = nnx.get_partition_spec(abstract_state) return jax.lax.with_sharding_constraint(abstract_state, sharding_specs) with mesh: # a jax.sharding.Mesh, covered fully in Week 10 abstract_target = jax.jit(create_abstract_sharded_target)() restored_sharded_state = mngr.restore( step, args=ocp.args.StandardRestore(abstract_target), ) __ __