File size: 586 Bytes
6759aa6 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 | 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)
|