from flax import nnx import chex import jax import jax.numpy as jnp class BulletproofBlock(nnx.Module): """A linear + ReLU block that validates its own inputs and outputs.""" def __init__(self, in_features: int, out_features: int, rngs: nnx.Rngs): self.in_features = in_features self.out_features = out_features self.linear = nnx.Linear(in_features, out_features, rngs=rngs) def __call__(self, x): # Input contract chex.assert_type(x, jnp.float32) chex.assert_rank(x, 2) chex.assert_axis_dimension(x, axis=1, expected=self.in_features) # The actual math x = self.linear(x) x = nnx.relu(x) # Output contract chex.assert_rank(x, 2) chex.assert_axis_dimension(x, axis=1, expected=self.out_features) chex.assert_tree_all_finite(x) return x __ __