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)