import jax import jax.numpy as jnp from flax import nnx import optax import orbax.checkpoint as ocp CHECKPOINT_DIR = "/tmp/resilient_training" SAVE_EVERY = 100 NUM_STEPS = 2000 def build_model_and_optimizer(seed=0): model = CNN(rngs=nnx.Rngs(seed)) tx = optax.adamw(learning_rate=0.001) optimizer = nnx.Optimizer(model, tx=tx, wrt=nnx.Param) return model, optimizer @nnx.jit def train_step(model, optimizer, batch): def loss_fn(model): logits = model(batch["image"]) loss = optax.softmax_cross_entropy_with_integer_labels( logits=logits, labels=batch["label"] ).mean() return loss, logits (loss, logits), grads = nnx.value_and_grad(loss_fn, has_aux=True)(model) optimizer.update(model, grads) return loss def main(): mngr = ocp.CheckpointManager( CHECKPOINT_DIR, options=ocp.CheckpointManagerOptions(max_to_keep=3, save_interval_steps=SAVE_EVERY), ) model, optimizer = build_model_and_optimizer(seed=0) start_step = 0 # Attempt to resume latest_step = mngr.latest_step() if latest_step is not None: print(f"Found existing checkpoint at step {latest_step}. Resuming.") abs_model = nnx.eval_shape(lambda: build_model_and_optimizer(seed=0)[0]) abs_optimizer = nnx.eval_shape( lambda: build_model_and_optimizer(seed=0)[1] ) graphdef, abs_params_state = nnx.split(abs_model, nnx.Param) abs_optimizer_state = nnx.state(abs_optimizer) restored = mngr.restore( latest_step, args=ocp.args.Composite( params=ocp.args.StandardRestore(abs_params_state), optimizer=ocp.args.StandardRestore(abs_optimizer_state), ), ) nnx.update(model, restored["params"]) nnx.update(optimizer, restored["optimizer"]) start_step = latest_step + 1 else: print("No checkpoint found. Starting fresh.") # Training loop for step in range(start_step, NUM_STEPS): # In a real project this batch would come from the Grain # pipeline built in Week 8 dummy_batch = { "image": jnp.ones((32, 28, 28, 1)), "label": jnp.zeros((32,), dtype=jnp.int32), } loss = train_step(model, optimizer, dummy_batch) if step % SAVE_EVERY == 0: params_state = nnx.split(optimizer, nnx.Param)[1] optimizer_state = nnx.state(optimizer) mngr.save( step, args=ocp.args.Composite( params=ocp.args.StandardSave(params_state), optimizer=ocp.args.StandardSave(optimizer_state), ), ) print(f"Step {step:4d} | Loss: {loss:.4f} | Checkpoint saved") mngr.wait_until_finished() mngr.close() print("Training complete.") if __name__ == "__main__": main() __ __