# Create a random number generator rngs = nnx.Rngs(42) # Seed for reproducibility # Create the model model = SimpleMLP(hidden_dim=128, output_dim=10, rngs=rngs) # Test with dummy data x = jnp.ones((32, 784)) # Batch of 32, 784 features each output = model(x) print(f"Input shape: {x.shape}") # (32, 784) print(f"Output shape: {output.shape}") # (32, 10) __ __