from flax.core import freeze, unfreeze # Set up PRNG key key = jax.random.PRNGKey(0) # Define input shape (batch_size, height, width, channels) input_shape = (1, 32, 32, 3) # Example for a 32x32 RGB image # Initialize model model = CNN(num_classes=10) # Assuming 10 output classes params = model.init(key, jnp.ones(input_shape))["params"] __ __