import jax from flax import linen as nn Array = jax.Array class Classifier(nn.Module): num_classes: int base_channels: int = 16 @nn.compact def __call__(self, x: Array) -> Array: h = nn.Conv(self.base_channels, kernel_size=(3, 3))(x) h = nn.relu(h) h = nn.max_pool(h, window_shape=(2, 2), strides=(2, 2)) h = nn.Conv(self.base_channels * 2, kernel_size=(3, 3))(h) h = nn.relu(h) h = nn.max_pool(h, window_shape=(2, 2), strides=(2, 2)) h = h.reshape(h.shape[0], -1) return nn.Dense(self.num_classes)(h)