from flax import nnx import jax.numpy as jnp from functools import partial class CNN(nnx.Module): """A simple CNN for image classification.""" def __init__(self, num_classes: int, *, rngs: nnx.Rngs): # Convolutional layers self.conv1 = nnx.Conv( in_features=1, # Input channels (1 for grayscale) out_features=32, # Output channels kernel_size=(3, 3), # 3x3 filters rngs=rngs ) self.conv2 = nnx.Conv( in_features=32, out_features=64, kernel_size=(3, 3), rngs=rngs ) # Dense layers # After two 2x2 pooling operations on 28x28 input: 28 → 14 → 7 # So we have 64 channels × 7 × 7 = 3136 features self.linear1 = nnx.Linear(3136, 256, rngs=rngs) self.linear2 = nnx.Linear(256, num_classes, rngs=rngs) # Pooling as a reusable operation self.pool = partial(nnx.avg_pool, window_shape=(2, 2), strides=(2, 2)) def __call__(self, x): # Block 1: Conv → ReLU → Pool x = self.conv1(x) x = nnx.relu(x) x = self.pool(x) # Block 2: Conv → ReLU → Pool x = self.conv2(x) x = nnx.relu(x) x = self.pool(x) # Flatten: (batch, height, width, channels) → (batch, features) x = x.reshape(x.shape[0], -1) # Classification head x = nnx.relu(self.linear1(x)) x = self.linear2(x) return x __ __