import jax import jax.numpy as jnp import chex @jax.jit def process_transactions(features, user_ids): # Type checks chex.assert_type(features, jnp.float32) chex.assert_type(user_ids, jnp.int32) # Rank checks chex.assert_rank(features, 2) # (batch, num_features) chex.assert_rank(user_ids, 1) # (batch,) # Shape checks: batch dimension is flexible, feature count is fixed chex.assert_shape(features, (None, 12)) # Ensure features and user_ids have matching batch sizes chex.assert_equal_shape_prefix([features, user_ids], prefix_len=1) return features * 0.1 # Placeholder processing __ __