import jax.numpy as jnp import chex x = jnp.ones((32, 12)) # Exact shape check chex.assert_shape(x, (32, 12)) # Use None for dimensions that can vary (e.g., batch size) chex.assert_shape(x, (None, 12)) # This will raise an AssertionError with a clear message chex.assert_shape(x, (32, 10)) __ __