| import jax | |
| from flax import linen as nn | |
| Array = jax.Array | |
| class Classifier(nn.Module): | |
| num_classes: int | |
| base_channels: int = 16 | |
| 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) | |