from flax import nnx import jax.numpy as jnp # Define a Model class MyModel(nnx.Module): def __init__(self, *, rngs: nnx.Rngs): self.linear = nnx.Linear(10, 5, rngs=rngs) def __call__(self, x): return nnx.relu(self.linear(x)) # Instantiate model = MyModel(rngs=nnx.Rngs(0)) # Forward Pass output = model(jnp.ones((32, 10))) # Inspect nnx.display(model) # Access Parameters weights = model.linear.kernel.value state = nnx.state(model, nnx.Param) # JIT Compile @nnx.jit def forward(model, x): return model(x) # Common Layers nnx.Linear(in_features, out_features, rngs=rngs) nnx.Conv(in_features, out_features, kernel_size, rngs=rngs) nnx.BatchNorm(num_features, rngs=rngs) nnx.Dropout(rate, rngs=rngs) nnx.Embed(num_embeddings, features, rngs=rngs) __ __