# Initialize model = CNN(num_classes=10, rngs=nnx.Rngs(0)) # Create dummy input: batch of 4 grayscale 28×28 images # Shape: (batch, height, width, channels) dummy_input = jnp.ones((4, 28, 28, 1)) # Forward pass output = model(dummy_input) print(f"Input shape: {dummy_input.shape}") # (4, 28, 28, 1) print(f"Output shape: {output.shape}") # (4, 10) # Inspect the model structure nnx.display(model) __ __