from flax import nnx class CNN(nnx.Module): def __init__(self, *, rngs: nnx.Rngs): self.conv1 = nnx.Conv(1, 32, kernel_size=(3, 3), rngs=rngs) self.conv2 = nnx.Conv(32, 64, kernel_size=(3, 3), rngs=rngs) self.linear = nnx.Linear(3136, 10, rngs=rngs) def __call__(self, x): x = nnx.relu(self.conv1(x)) x = nnx.relu(self.conv2(x)) x = x.reshape(x.shape[0], -1) return self.linear(x) model = CNN(rngs=nnx.Rngs(0)) nnx.display(model) __ __