from flax import nnx import jax.numpy as jnp class SimpleMLP(nnx.Module): def __init__(self, hidden_dim: int, output_dim: int, *, rngs: nnx.Rngs): # Define layers as attributes self.linear1 = nnx.Linear(784, hidden_dim, rngs=rngs) self.linear2 = nnx.Linear(hidden_dim, output_dim, rngs=rngs) def __call__(self, x): # Define forward pass x = self.linear1(x) x = nnx.relu(x) x = self.linear2(x) return x __ __