import chex import jax.numpy as jnp # Shape and type contracts chex.assert_shape(x, (None, 12)) # None = any size allowed chex.assert_rank(x, 2) # Exactly 2 dimensions chex.assert_type(x, jnp.float32) # Enforce dtype chex.assert_axis_dimension(x, axis=1, expected=12) chex.assert_equal_shape_prefix([a, b], prefix_len=1) # Matching batch dims # NaN/Inf Hunting chex.assert_tree_all_finite(pytree) # Works on nested structures # Runtime value assertions @chex.chexify @jax.jit def fn(x): chex.assert_equal(jnp.all(x >= 0), True) return x chex.block_until_chexify_assertions_complete() # Recompilation detection @chex.assert_max_traces(n=2) @jax.jit def fn(x): return x chex.clear_trace_counter() # Reset between tests # Comparing PyTrees chex.assert_trees_all_close(tree1, tree2, rtol=1e-5) chex.assert_trees_all_equal(tree1, tree2) __ __