# Initialize model model = CNN(rngs=nnx.Rngs(0)) # Test forward pass dummy_input = jnp.ones((1, 28, 28, 1)) dummy_output = model(dummy_input) print(f"Model output shape: {dummy_output.shape}") # (1, 10) # Initialize optimizer learning_rate = 0.005 momentum = 0.9 tx = optax.adamw(learning_rate, momentum) optimizer = nnx.Optimizer(model, tx, wrt=nnx.Param) # Display model structure nnx.display(model) __ __