from flax import nnx class NNX_MLP(nnx.Module): def __init__(self, rngs: nnx.Rngs): self.linear1 = nnx.Linear(784, 128, rngs=rngs) self.linear2 = nnx.Linear(128, 10, rngs=rngs) def __call__(self, x): x = nnx.relu(self.linear1(x)) return self.linear2(x) model = NNX_MLP(rngs=nnx.Rngs(0)) output = model(input_data) # Looks just like PyTorch __ __