import jax.numpy as jnp import flax.linen as nn class CNN(nn.Module): num_classes: int # Number of output classes @nn.compact def __call__(self, x): x = nn.Conv(features=32, kernel_size=(3, 3), strides=(1, 1), padding="SAME")(x) x = nn.relu(x) x = nn.max_pool(x, window_shape=(2, 2), strides=(2, 2)) x = nn.Conv(features=64, kernel_size=(3, 3), strides=(1, 1), padding="SAME")(x) x = nn.relu(x) x = nn.max_pool(x, window_shape=(2, 2), strides=(2, 2)) x = nn.Conv(features=128, kernel_size=(3, 3), strides=(1, 1), padding="SAME")(x) x = nn.relu(x) x = nn.max_pool(x, window_shape=(2, 2), strides=(2, 2)) x = x.reshape((x.shape[0], -1)) # Flatten feature maps x = nn.Dense(features=128)(x) x = nn.relu(x) x = nn.Dense(features=self.num_classes)(x) # Output layer return x __ __