from flax import nnx model = CNN(rngs=nnx.Rngs(0)) # Extract the model's state as a pytree state = nnx.state(model) __ __