rngs = nnx.Rngs(jax.random.PRNGKey(0)) safe_layer = BulletproofBlock(in_features=64, out_features=128, rngs=rngs) # Correct input: should pass good_data = jnp.ones((32, 64)) output = safe_layer(good_data) print(f"Success! Output shape: {output.shape}") # Wrong input: should fail with a clear error bad_data = jnp.ones((32, 32)) # Wrong feature dimension try: output = safe_layer(bad_data) except Exception as e: print(f"Caught the expected error:\n{e}") __ __